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>
This commit is contained in:
parent
695f8e84f7
commit
0d22946a0e
1835 changed files with 170291 additions and 2849 deletions
|
|
@ -47,6 +47,7 @@
|
|||
#include "portable-snippets/exact-int/exact-int.h"
|
||||
|
||||
#include "dpf/utils.hpp"
|
||||
#include "dpf/twiddle.hpp"
|
||||
#include "dpf/bit_array.hpp"
|
||||
|
||||
namespace dpf
|
||||
|
|
@ -55,16 +56,6 @@ namespace dpf
|
|||
namespace detail
|
||||
{
|
||||
|
||||
template <typename Iterator>
|
||||
struct extract_bit_simde_node
|
||||
{
|
||||
bool operator()(Iterator it) const
|
||||
{
|
||||
auto buf = reinterpret_cast<const char *>(&*it);
|
||||
return buf[0] & 1;
|
||||
}
|
||||
};
|
||||
|
||||
template <typename NodeT, typename Iterator>
|
||||
struct extract_bit;
|
||||
|
||||
|
|
@ -72,11 +63,23 @@ HEDLEY_PRAGMA(GCC diagnostic push)
|
|||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
||||
template <typename Iterator>
|
||||
struct extract_bit<simde__m128i, Iterator>
|
||||
: public extract_bit_simde_node<Iterator> { };
|
||||
{
|
||||
bool operator()(Iterator it) const
|
||||
{
|
||||
// Same lo-bit extract used by tree walk / correction packing.
|
||||
return static_cast<bool>(dpf::get_lo_bit(*it));
|
||||
}
|
||||
};
|
||||
|
||||
template <typename Iterator>
|
||||
struct extract_bit<simde__m256i, Iterator>
|
||||
: public extract_bit_simde_node<Iterator> { };
|
||||
{
|
||||
bool operator()(Iterator it) const
|
||||
{
|
||||
auto buf = reinterpret_cast<const char *>(&*it);
|
||||
return buf[0] & 1;
|
||||
}
|
||||
};
|
||||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||||
|
||||
} // namespace detail
|
||||
|
|
|
|||
161
include/dpf/aes_sbox_bp.hpp
Normal file
161
include/dpf/aes_sbox_bp.hpp
Normal file
|
|
@ -0,0 +1,161 @@
|
|||
/// @file dpf/aes_sbox_bp.hpp
|
||||
/// @brief Boyar–Peralta AES S-box (Yale CMT SLP, 32 ANDs) over XOR bit shares.
|
||||
/// @details Circuit: http://www.cs.yale.edu/homes/peralta/CircuitStuff/SLP_AES_113.txt
|
||||
/// (Boyar and Peralta, ePrint 2011/332). Matches the AES S-box used
|
||||
/// by `aes_ref` / hardware AES. `#` is XNOR.
|
||||
|
||||
#ifndef LIBDPF_INCLUDE_DPF_AES_SBOX_BP_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_AES_SBOX_BP_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
|
||||
namespace aes_bp
|
||||
{
|
||||
|
||||
inline constexpr std::size_t wire_count = 121;
|
||||
inline constexpr std::size_t op_count = 113;
|
||||
inline constexpr std::size_t and_count = 32;
|
||||
/// @brief Multiplicative depth of the SLP (XOR/XNOR are free).
|
||||
inline constexpr std::size_t and_layer_count = 6;
|
||||
inline constexpr std::uint8_t out_wire[8] = {114,117,119,107,116,120,115,111};
|
||||
|
||||
/// @brief AND-op indices in `ops` grouped by multiplicative depth.
|
||||
/// @details XOR/XNOR between layers stay local. Every AND in a layer has
|
||||
/// both inputs ready, so one exchange covers the whole layer
|
||||
/// (and every S-box byte sharing that schedule).
|
||||
inline constexpr std::uint8_t and_layer_size[and_layer_count] = {9, 1, 2, 7, 5, 8};
|
||||
inline constexpr std::uint8_t and_layer_ops[and_layer_count][9] = {
|
||||
{23, 24, 26, 28, 29, 31, 33, 34, 36},
|
||||
{47},
|
||||
{49, 53},
|
||||
{57, 69, 72, 73, 78, 81, 82},
|
||||
{60, 67, 68, 76, 77},
|
||||
{70, 71, 74, 75, 79, 80, 83, 84},
|
||||
};
|
||||
static_assert(
|
||||
and_layer_size[0] + and_layer_size[1] + and_layer_size[2]
|
||||
+ and_layer_size[3] + and_layer_size[4] + and_layer_size[5]
|
||||
== and_count,
|
||||
"Boyar–Peralta AND layer sizes must cover every AND");
|
||||
|
||||
// kind: 0=XOR, 1=AND, 2=XNOR. Each row is {kind, dst, a, b}.
|
||||
inline constexpr std::uint8_t ops[op_count][4] = {
|
||||
{0, 8, 3, 5},
|
||||
{0, 9, 0, 6},
|
||||
{0, 10, 0, 3},
|
||||
{0, 11, 0, 5},
|
||||
{0, 12, 1, 2},
|
||||
{0, 13, 12, 7},
|
||||
{0, 14, 13, 3},
|
||||
{0, 15, 9, 8},
|
||||
{0, 16, 13, 0},
|
||||
{0, 17, 13, 6},
|
||||
{0, 18, 17, 11},
|
||||
{0, 19, 4, 15},
|
||||
{0, 20, 19, 5},
|
||||
{0, 21, 19, 1},
|
||||
{0, 22, 20, 7},
|
||||
{0, 23, 20, 12},
|
||||
{0, 24, 21, 10},
|
||||
{0, 25, 7, 24},
|
||||
{0, 26, 23, 24},
|
||||
{0, 27, 23, 11},
|
||||
{0, 28, 12, 24},
|
||||
{0, 29, 9, 28},
|
||||
{0, 30, 0, 28},
|
||||
{1, 31, 15, 20},
|
||||
{1, 32, 18, 22},
|
||||
{0, 33, 32, 31},
|
||||
{1, 34, 14, 7},
|
||||
{0, 35, 34, 31},
|
||||
{1, 36, 9, 28},
|
||||
{1, 37, 17, 13},
|
||||
{0, 38, 37, 36},
|
||||
{1, 39, 16, 25},
|
||||
{0, 40, 39, 36},
|
||||
{1, 41, 10, 24},
|
||||
{1, 42, 8, 26},
|
||||
{0, 43, 42, 41},
|
||||
{1, 44, 11, 23},
|
||||
{0, 45, 44, 41},
|
||||
{0, 46, 33, 21},
|
||||
{0, 47, 35, 45},
|
||||
{0, 48, 38, 43},
|
||||
{0, 49, 40, 45},
|
||||
{0, 50, 46, 43},
|
||||
{0, 51, 47, 27},
|
||||
{0, 52, 48, 29},
|
||||
{0, 53, 49, 30},
|
||||
{0, 54, 50, 51},
|
||||
{1, 55, 50, 52},
|
||||
{0, 56, 53, 55},
|
||||
{1, 57, 54, 56},
|
||||
{0, 58, 57, 51},
|
||||
{0, 59, 52, 53},
|
||||
{0, 60, 51, 55},
|
||||
{1, 61, 60, 59},
|
||||
{0, 62, 61, 53},
|
||||
{0, 63, 52, 62},
|
||||
{0, 64, 56, 62},
|
||||
{1, 65, 53, 64},
|
||||
{0, 66, 65, 63},
|
||||
{0, 67, 56, 65},
|
||||
{1, 68, 58, 67},
|
||||
{0, 69, 54, 68},
|
||||
{0, 70, 69, 66},
|
||||
{0, 71, 58, 62},
|
||||
{0, 72, 58, 69},
|
||||
{0, 73, 62, 66},
|
||||
{0, 74, 71, 70},
|
||||
{1, 75, 73, 20},
|
||||
{1, 76, 66, 22},
|
||||
{1, 77, 62, 7},
|
||||
{1, 78, 72, 28},
|
||||
{1, 79, 69, 13},
|
||||
{1, 80, 58, 25},
|
||||
{1, 81, 71, 24},
|
||||
{1, 82, 74, 26},
|
||||
{1, 83, 70, 23},
|
||||
{1, 84, 73, 15},
|
||||
{1, 85, 66, 18},
|
||||
{1, 86, 62, 14},
|
||||
{1, 87, 72, 9},
|
||||
{1, 88, 69, 17},
|
||||
{1, 89, 58, 16},
|
||||
{1, 90, 71, 10},
|
||||
{1, 91, 74, 8},
|
||||
{1, 92, 70, 11},
|
||||
{0, 93, 90, 91},
|
||||
{0, 94, 85, 93},
|
||||
{0, 95, 84, 94},
|
||||
{0, 96, 75, 77},
|
||||
{0, 97, 76, 75},
|
||||
{0, 98, 78, 79},
|
||||
{0, 99, 87, 96},
|
||||
{0, 100, 82, 98},
|
||||
{0, 101, 83, 99},
|
||||
{0, 102, 100, 101},
|
||||
{0, 103, 98, 97},
|
||||
{0, 104, 78, 80},
|
||||
{0, 105, 88, 93},
|
||||
{0, 106, 96, 104},
|
||||
{0, 107, 95, 103},
|
||||
{0, 108, 81, 100},
|
||||
{0, 109, 89, 102},
|
||||
{0, 110, 105, 106},
|
||||
{2, 111, 87, 110},
|
||||
{0, 112, 90, 108},
|
||||
{0, 113, 94, 86},
|
||||
{0, 114, 95, 108},
|
||||
{2, 115, 102, 110},
|
||||
{0, 116, 106, 107},
|
||||
{2, 117, 107, 108},
|
||||
{0, 118, 109, 112},
|
||||
{2, 119, 118, 92},
|
||||
{0, 120, 113, 109},
|
||||
};
|
||||
|
||||
} // namespace aes_bp
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_AES_SBOX_BP_HPP__
|
||||
633
include/dpf/app_flow.hpp
Normal file
633
include/dpf/app_flow.hpp
Normal file
|
|
@ -0,0 +1,633 @@
|
|||
/// @file dpf/app_flow.hpp
|
||||
/// @brief Measure an application plan on an explicit `run_config`.
|
||||
/// @details `exercise_plan(p, ex, cfg)` drives both parties of `p` over
|
||||
/// `cfg.kind` (in-process memory, sync streams, async memory, unix
|
||||
/// sockets, TCP mux, parallel TCP, or SCTP) with the configured lanes,
|
||||
/// framing, instances, window, chunking, pipelining, compute threads,
|
||||
/// warmup, and trials. Link setup is outside the timed region. The
|
||||
/// overloads without a config read `run_config::from_env()` once, for
|
||||
/// the example binaries.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_APP_FLOW_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_APP_FLOW_HPP__
|
||||
|
||||
#include <algorithm>
|
||||
#include <atomic>
|
||||
#include <chrono>
|
||||
#include <cmath>
|
||||
#include <condition_variable>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstdlib>
|
||||
#include <cstring>
|
||||
#include <exception>
|
||||
#include <iostream>
|
||||
#include <map>
|
||||
#include <mutex>
|
||||
#include <optional>
|
||||
#include <string>
|
||||
#include <thread>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/net/asio_ns.hpp"
|
||||
|
||||
#include "dpf/app_runtime.hpp"
|
||||
#include "dpf/compose.hpp"
|
||||
#include "dpf/compose_async.hpp"
|
||||
#include "dpf/experiment.hpp"
|
||||
#include "dpf/net/async_round_sink.hpp"
|
||||
#include "dpf/net/memory_sink.hpp"
|
||||
#include "dpf/net/round_lane.hpp"
|
||||
#include "dpf/net/stream_array.hpp"
|
||||
#include "dpf/online_session.hpp"
|
||||
#include "dpf/party_runner.hpp"
|
||||
#include "dpf/run_config.hpp"
|
||||
#include "dpf/run_log.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace app
|
||||
{
|
||||
|
||||
using transport_kind = net::transport;
|
||||
|
||||
inline const char * transport_name(transport_kind t) noexcept
|
||||
{
|
||||
return net::transport_name(t);
|
||||
}
|
||||
|
||||
/// @brief `DPF_TRANSPORT` (default `async`). Unknown names throw.
|
||||
inline transport_kind transport_from_env()
|
||||
{
|
||||
return run_config::from_env().kind;
|
||||
}
|
||||
|
||||
/// @brief Drive a plan through `schedule_session` on an `async_round_sink`.
|
||||
inline void drive_via_schedule_async(const protocol::plan & p,
|
||||
net::async_round_sink & sink,
|
||||
std::vector<std::vector<std::uint8_t>> & values,
|
||||
const std::map<std::uint32_t, protocol::kernel_fn> & kernels,
|
||||
std::size_t party, const protocol::drive_options & opt = {})
|
||||
{
|
||||
protocol::drive_via_schedule(p, sink, values, kernels, party, opt);
|
||||
}
|
||||
|
||||
/// @brief What one `exercise_plan` measured.
|
||||
/// @details `bytes` is the plan's slot bytes (one instance). `wire_*` are
|
||||
/// party 0's link counters including every header. `wall_ns` is the
|
||||
/// median party-0 drive time over the timed trials.
|
||||
struct cost
|
||||
{
|
||||
std::size_t rounds = 0;
|
||||
std::size_t bytes = 0;
|
||||
std::uint64_t wire_out = 0;
|
||||
std::uint64_t wire_in = 0;
|
||||
std::uint64_t frames_out = 0;
|
||||
std::uint64_t frames_in = 0;
|
||||
std::uint64_t write_calls = 0;
|
||||
std::uint64_t wall_ns = 0;
|
||||
};
|
||||
|
||||
inline cost plan_cost(const protocol::plan & p)
|
||||
{
|
||||
cost out;
|
||||
out.rounds = p.rounds();
|
||||
for (auto n : p.slot_bytes_all())
|
||||
out.bytes += n;
|
||||
return out;
|
||||
}
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
struct once
|
||||
{
|
||||
std::uint64_t wall_ns = 0;
|
||||
net::stream_stats wire;
|
||||
/// Every party's drive time (`party_walls[0] == wall_ns`).
|
||||
std::vector<std::uint64_t> party_walls;
|
||||
};
|
||||
|
||||
/// @brief Parties 0 and 1 on two threads over a paired in-process sink. Party
|
||||
/// `p` draws from `seeds->derive_party(p)` when `seeds` is set; `ex`
|
||||
/// records party 0 and receives both parties' noted seeds.
|
||||
template <typename Drive>
|
||||
inline once two_threads(experiment * ex, const experiment * seeds, Drive drive)
|
||||
{
|
||||
once out;
|
||||
out.party_walls.assign(2, 0);
|
||||
std::exception_ptr err[2];
|
||||
std::vector<experiment::noted_seed> noted[2];
|
||||
start_gate gate(2);
|
||||
auto side = [&](unsigned party) {
|
||||
try
|
||||
{
|
||||
std::optional<experiment> stream;
|
||||
if (seeds != nullptr)
|
||||
stream.emplace(seeds->derive_party(party));
|
||||
if (!gate.arrive_and_wait())
|
||||
throw gate_broken();
|
||||
experiment * mine = party == 0 ? ex : nullptr;
|
||||
protocol::round_probe probe{};
|
||||
if (mine != nullptr)
|
||||
{
|
||||
probe = mine->probe();
|
||||
mine->begin_timing();
|
||||
}
|
||||
const auto a = std::chrono::steady_clock::now();
|
||||
drive(party, mine != nullptr ? &probe : nullptr);
|
||||
out.party_walls[party] = static_cast<std::uint64_t>(
|
||||
std::chrono::duration_cast<std::chrono::nanoseconds>(
|
||||
std::chrono::steady_clock::now() - a)
|
||||
.count());
|
||||
if (mine != nullptr)
|
||||
mine->end_timing();
|
||||
if (stream)
|
||||
noted[party] = stream->seeds();
|
||||
}
|
||||
catch (...)
|
||||
{
|
||||
err[party] = std::current_exception();
|
||||
gate.fail();
|
||||
}
|
||||
};
|
||||
std::thread t0(side, 0u);
|
||||
std::thread t1(side, 1u);
|
||||
t0.join();
|
||||
t1.join();
|
||||
for (auto & e : err)
|
||||
{
|
||||
if (!e)
|
||||
continue;
|
||||
try
|
||||
{
|
||||
std::rethrow_exception(e);
|
||||
}
|
||||
catch (const gate_broken &)
|
||||
{
|
||||
continue;
|
||||
}
|
||||
catch (...)
|
||||
{
|
||||
}
|
||||
std::rethrow_exception(e);
|
||||
}
|
||||
out.wall_ns = out.party_walls[0];
|
||||
if (ex != nullptr && seeds != nullptr)
|
||||
for (unsigned p = 0; p < 2; ++p)
|
||||
ex->fold_seeds("p" + std::to_string(p), noted[p]);
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief One run of `plans` on `cfg.kind` from fresh copies of `inputs`;
|
||||
/// `ex` records party 0 when set, and parties draw from `seeds`'s
|
||||
/// derived streams when it is set.
|
||||
inline once run_once(const std::vector<protocol::plan> & plans,
|
||||
const std::vector<party_values> & inputs,
|
||||
const std::map<std::uint32_t, protocol::kernel_fn> & kernels,
|
||||
const run_config & cfg, experiment * ex, const experiment * seeds = nullptr)
|
||||
{
|
||||
std::vector<party_values> values = inputs;
|
||||
values.resize(plans.size());
|
||||
const auto slots = plans[0].slot_bytes_all();
|
||||
auto opt = cfg.drive();
|
||||
const bool in_memory = cfg.kind == net::transport::memory_sink
|
||||
|| cfg.kind == net::transport::memory_stream;
|
||||
if (in_memory && plans.size() != 2)
|
||||
throw std::invalid_argument(std::string("transport ")
|
||||
+ net::transport_name(cfg.kind) + " runs two parties");
|
||||
switch (cfg.kind)
|
||||
{
|
||||
case net::transport::memory_sink:
|
||||
{
|
||||
auto sl = slots.empty() ? std::vector<std::size_t>{0} : slots;
|
||||
auto sinks = net::make_memory_sink_pair(cfg.instances, sl);
|
||||
return two_threads(ex, seeds, [&](unsigned party, const protocol::round_probe * pr) {
|
||||
auto o = opt;
|
||||
o.probe = pr;
|
||||
protocol::drive_via_schedule(plans[party],
|
||||
party == 0 ? sinks.first : sinks.second, values[party], kernels,
|
||||
party, o);
|
||||
});
|
||||
}
|
||||
case net::transport::memory_stream:
|
||||
{
|
||||
const std::size_t nstreams =
|
||||
slots.empty() ? 1 : protocol::lanes_for_plan(slots.size(), opt);
|
||||
auto ends = net::make_memory_stream_pair(nstreams);
|
||||
return two_threads(ex, seeds, [&](unsigned party, const protocol::round_probe * pr) {
|
||||
auto o = opt;
|
||||
o.probe = pr;
|
||||
protocol::drive_plan_on_streams(plans[party],
|
||||
party == 0 ? ends.first : ends.second, values[party], kernels,
|
||||
party, cfg.instances, o);
|
||||
});
|
||||
}
|
||||
default:
|
||||
{
|
||||
const auto r = run_parties(plans, values, kernels, cfg, ex, nullptr, seeds);
|
||||
once out;
|
||||
out.wall_ns = r.party0_wall_ns;
|
||||
out.wire = r.wire[0];
|
||||
out.party_walls = r.party_wall_ns;
|
||||
return out;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/// @brief Drive every party's plan over `cfg` and return party 0's cost.
|
||||
/// @details `plans[i]` is party `i`'s plan (2 or 3 parties); each trial starts
|
||||
/// from a fresh copy of `inputs` (missing entries start empty).
|
||||
/// Warmup runs are not timed; `wall_ns` is the median timed trial.
|
||||
inline cost exercise_parties(const std::vector<protocol::plan> & plans,
|
||||
const std::vector<party_values> & inputs,
|
||||
const std::map<std::uint32_t, protocol::kernel_fn> & kernels, experiment * ex,
|
||||
const run_config & cfg)
|
||||
{
|
||||
if (plans.size() < 2 || plans.size() > 3)
|
||||
throw std::invalid_argument("exercise_parties needs 2 or 3 plans");
|
||||
const protocol::plan & p = plans[0];
|
||||
cost out = plan_cost(p);
|
||||
if (ex != nullptr)
|
||||
{
|
||||
ex->ingest_plan(p);
|
||||
ex->set_config(cfg.describe());
|
||||
}
|
||||
if (p.slot_bytes_all().empty() && cfg.kind != net::transport::memory_sink)
|
||||
{
|
||||
DPF_LOG(warning, "trials.skipped")
|
||||
.kv("experiment", ex != nullptr ? ex->name() : std::string("none"))
|
||||
.kv("detail", "the plan has no exchange rounds; nothing was timed and "
|
||||
"wall_ns stays 0");
|
||||
return out;
|
||||
}
|
||||
const std::size_t total = cfg.warmup + std::max<std::size_t>(1, cfg.trials);
|
||||
std::vector<std::uint64_t> walls;
|
||||
std::vector<std::uint64_t> slowest;
|
||||
detail::once last;
|
||||
for (std::size_t t = 0; t < total; ++t)
|
||||
{
|
||||
const bool timed = t >= cfg.warmup;
|
||||
const bool record = t + 1 == total;
|
||||
last = detail::run_once(plans, inputs, kernels, cfg, record ? ex : nullptr, ex);
|
||||
DPF_LOG(debug, "trial").kv("index", t).kv("timed", timed)
|
||||
.kv("instrumented", record && ex != nullptr).kv("p0_wall_ns", last.wall_ns)
|
||||
.kv("wire_out", last.wire.bytes_out).kv("wire_in", last.wire.bytes_in);
|
||||
if (timed)
|
||||
{
|
||||
walls.push_back(last.wall_ns);
|
||||
std::uint64_t w = last.wall_ns;
|
||||
for (auto p : last.party_walls)
|
||||
w = std::max(w, p);
|
||||
slowest.push_back(w);
|
||||
if (ex != nullptr)
|
||||
ex->add_trial(last.wall_ns, last.party_walls);
|
||||
}
|
||||
}
|
||||
std::sort(walls.begin(), walls.end());
|
||||
std::sort(slowest.begin(), slowest.end());
|
||||
out.wall_ns = walls.empty() ? 0 : walls[walls.size() / 2];
|
||||
if (!walls.empty() && log::enabled(log::level::info))
|
||||
{
|
||||
double mean = 0;
|
||||
for (auto w : walls)
|
||||
mean += static_cast<double>(w);
|
||||
mean /= static_cast<double>(walls.size());
|
||||
double var = 0;
|
||||
for (auto w : walls)
|
||||
var += (static_cast<double>(w) - mean) * (static_cast<double>(w) - mean);
|
||||
const double sd = walls.size() > 1
|
||||
? std::sqrt(var / static_cast<double>(walls.size() - 1)) : 0.0;
|
||||
DPF_LOG(info, "trials")
|
||||
.kv("experiment", ex != nullptr ? ex->name() : std::string("none"))
|
||||
.kv("transport", net::transport_name(cfg.kind)).kv("n", walls.size())
|
||||
.kv("warmup", cfg.warmup).kv("median_ns", out.wall_ns)
|
||||
.kv("min_ns", walls.front()).kv("max_ns", walls.back())
|
||||
.kv("mean_ns", mean).kv("sd_ns", sd)
|
||||
.kv("slowest_median_ns", slowest[slowest.size() / 2])
|
||||
.kv("instrumented_trial", "last")
|
||||
.kv("links", "rebuilt per trial");
|
||||
}
|
||||
out.wire_out = last.wire.bytes_out;
|
||||
out.wire_in = last.wire.bytes_in;
|
||||
out.frames_out = last.wire.frames_out;
|
||||
out.frames_in = last.wire.frames_in;
|
||||
out.write_calls = last.wire.write_calls;
|
||||
if (ex != nullptr)
|
||||
{
|
||||
experiment::wire_counts w;
|
||||
w.bytes_out = last.wire.bytes_out;
|
||||
w.bytes_in = last.wire.bytes_in;
|
||||
w.payload_out = last.wire.payload_out;
|
||||
w.payload_in = last.wire.payload_in;
|
||||
w.frames_out = last.wire.frames_out;
|
||||
w.frames_in = last.wire.frames_in;
|
||||
w.write_calls = last.wire.write_calls;
|
||||
ex->set_wire(w);
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief Drive `p` as both parties' plan over `cfg` and return its cost.
|
||||
inline cost exercise_plan(const protocol::plan & p, experiment * ex,
|
||||
const run_config & cfg,
|
||||
const std::map<std::uint32_t, protocol::kernel_fn> & kernels = {})
|
||||
{
|
||||
return exercise_parties({p, p}, {}, kernels, ex, cfg);
|
||||
}
|
||||
|
||||
/// @brief `exercise_plan` on `run_config::from_env()`.
|
||||
inline cost exercise_plan(const protocol::plan & p, experiment * ex = nullptr)
|
||||
{
|
||||
return exercise_plan(p, ex, run_config::from_env());
|
||||
}
|
||||
|
||||
inline experiment measure_plan(const char * name, const protocol::plan & p,
|
||||
const run_config & cfg)
|
||||
{
|
||||
experiment ex(name, "p0");
|
||||
(void)exercise_plan(p, &ex, cfg);
|
||||
return ex;
|
||||
}
|
||||
|
||||
inline experiment measure_plan(const char * name, const protocol::plan & p)
|
||||
{
|
||||
return measure_plan(name, p, run_config::from_env());
|
||||
}
|
||||
|
||||
/// @brief Schedule `c`, then `exercise_plan`.
|
||||
inline cost exercise(protocol::composer & c)
|
||||
{
|
||||
return exercise_plan(c.default_plan());
|
||||
}
|
||||
|
||||
/// @brief Exercise `p` and require `expect_rounds`. Prints `name rounds= bytes=`.
|
||||
inline int run_plan(const char * name, const protocol::plan & p,
|
||||
std::size_t expect_rounds, const run_config & cfg = run_config::from_env())
|
||||
{
|
||||
try
|
||||
{
|
||||
start_logging(cfg);
|
||||
const cost got = exercise_plan(p, nullptr, cfg);
|
||||
std::cout << name << " rounds=" << got.rounds << " bytes=" << got.bytes
|
||||
<< " wire_out=" << got.wire_out << " wire_in=" << got.wire_in
|
||||
<< "\n";
|
||||
if (got.rounds != expect_rounds)
|
||||
{
|
||||
std::cerr << name << " expected " << expect_rounds << " rounds\n";
|
||||
return 1;
|
||||
}
|
||||
}
|
||||
catch (const std::exception & ex)
|
||||
{
|
||||
DPF_LOG(error, "run.failed").kv("name", name).kv("what", ex.what());
|
||||
std::cerr << name << " flow: " << ex.what() << "\n";
|
||||
return 1;
|
||||
}
|
||||
return 0;
|
||||
}
|
||||
|
||||
/// @brief Measure `p`, print cost, and write CSVs under `DPF_EXPERIMENT_DIR`.
|
||||
/// @details When `prep` is set, the prep for that demand is dealt and shipped
|
||||
/// over `cfg`'s transport and its size is printed.
|
||||
inline int run_measured(const char * name, const protocol::plan & p,
|
||||
std::size_t expect_rounds, const run_config & cfg = run_config::from_env(),
|
||||
const prep::demand * prep = nullptr)
|
||||
{
|
||||
try
|
||||
{
|
||||
start_logging(cfg);
|
||||
auto ex = measure_plan(name, p, cfg);
|
||||
std::cout << name << " transport=" << net::transport_name(cfg.kind)
|
||||
<< " rounds=" << ex.interactive_rounds()
|
||||
<< " bytes=" << ex.plan_bytes_out()
|
||||
<< " wire_out=" << ex.wire().bytes_out
|
||||
<< " wire_in=" << ex.wire().bytes_in
|
||||
<< " wall_ns=" << ex.wall_ns()
|
||||
<< " median_ns=" << ex.median_trial_ns()
|
||||
<< " cpu_ns=" << ex.cpu_ns()
|
||||
<< " prg_evals=" << ex.prg_evals()
|
||||
<< " random_bytes=" << ex.random_bytes()
|
||||
<< " seed=" << ex.seed_hex();
|
||||
if (prep != nullptr)
|
||||
{
|
||||
const auto shipped = session::ship_prep(*prep, cfg);
|
||||
std::cout << " prep_bytes=" << shipped.bytes0;
|
||||
}
|
||||
std::cout << "\n";
|
||||
if (ex.interactive_rounds() != expect_rounds)
|
||||
{
|
||||
std::cerr << name << " expected " << expect_rounds << " rounds\n";
|
||||
return 1;
|
||||
}
|
||||
if (const char * dir = std::getenv("DPF_EXPERIMENT_DIR"))
|
||||
{
|
||||
if (dir[0] != '\0')
|
||||
ex.write_csv(dir);
|
||||
}
|
||||
}
|
||||
catch (const std::exception & ex)
|
||||
{
|
||||
DPF_LOG(error, "run.failed").kv("name", name).kv("what", ex.what());
|
||||
std::cerr << name << " measure: " << ex.what() << "\n";
|
||||
return 1;
|
||||
}
|
||||
return 0;
|
||||
}
|
||||
|
||||
/// @brief Exercise `c` and require `expect_rounds`.
|
||||
inline int run(const char * name, protocol::composer & c, std::size_t expect_rounds)
|
||||
{
|
||||
return run_plan(name, c.default_plan(), expect_rounds);
|
||||
}
|
||||
|
||||
/// @brief `run_measured` on `c.default_plan()`.
|
||||
inline int run_measured(const char * name, protocol::composer & c,
|
||||
std::size_t expect_rounds)
|
||||
{
|
||||
return run_measured(name, c.default_plan(), expect_rounds);
|
||||
}
|
||||
|
||||
/// @brief Many instances, both parties, delays that do not line up.
|
||||
/// @details A worker never sits on a side whose peer is the one that still
|
||||
/// has to submit. Parked receives yield the thread to whichever
|
||||
/// side is behind, so a slow step does not stall the rest and a
|
||||
/// round-robin that always resumes the waiter cannot deadlock.
|
||||
inline void run_fleet(protocol::composer & c, std::size_t instances,
|
||||
std::uint32_t chaos_seed)
|
||||
{
|
||||
if (instances == 0)
|
||||
throw std::invalid_argument("run_fleet needs instances");
|
||||
auto plan = c.default_plan();
|
||||
auto slots = plan.slot_bytes_all();
|
||||
if (slots.empty())
|
||||
slots.push_back(0);
|
||||
|
||||
struct side
|
||||
{
|
||||
protocol::drive_cursor cursor{};
|
||||
std::vector<std::vector<std::uint8_t>> values;
|
||||
net::memory_sink * sink = nullptr;
|
||||
std::chrono::steady_clock::time_point ready_at{};
|
||||
bool busy = false;
|
||||
};
|
||||
struct inst
|
||||
{
|
||||
std::pair<net::memory_sink, net::memory_sink> sinks;
|
||||
side party[2]{};
|
||||
};
|
||||
std::vector<inst> all;
|
||||
all.reserve(instances);
|
||||
auto chaos_us = [&](std::size_t id, int party, std::size_t step) {
|
||||
std::uint32_t x = chaos_seed
|
||||
^ static_cast<std::uint32_t>(id * 0x9E3779B9u)
|
||||
^ static_cast<std::uint32_t>(party * 0x85EBCA6Bu)
|
||||
^ static_cast<std::uint32_t>(step * 0xC2B2AE35u);
|
||||
x ^= x << 13;
|
||||
x ^= x >> 17;
|
||||
// Mostly short, with occasional long stalls. Not sorted by instance.
|
||||
const std::uint32_t bucket = x % 17u;
|
||||
if (party == 0 && step == 0 && (id % 5u) == 0)
|
||||
return 12000;
|
||||
if (bucket == 0)
|
||||
return 1500;
|
||||
if (bucket < 4)
|
||||
return static_cast<int>(x % 400u);
|
||||
return static_cast<int>(x % 40u);
|
||||
};
|
||||
for (std::size_t i = 0; i < instances; ++i)
|
||||
{
|
||||
all.push_back(inst{net::make_memory_sink_pair(1, slots), {}});
|
||||
auto & created = all.back();
|
||||
created.party[0].sink = &created.sinks.first;
|
||||
created.party[1].sink = &created.sinks.second;
|
||||
const auto now = std::chrono::steady_clock::now();
|
||||
created.party[0].ready_at = now + std::chrono::microseconds(
|
||||
chaos_us(i, 0, 0));
|
||||
created.party[1].ready_at = now + std::chrono::microseconds(
|
||||
chaos_us(i, 1, 0));
|
||||
}
|
||||
|
||||
std::mutex mu;
|
||||
std::condition_variable cv;
|
||||
std::map<std::uint32_t, protocol::kernel_fn> kernels;
|
||||
std::size_t finished = 0;
|
||||
const std::size_t workers = std::min<std::size_t>(
|
||||
std::thread::hardware_concurrency() == 0
|
||||
? 4
|
||||
: std::thread::hardware_concurrency(),
|
||||
instances);
|
||||
|
||||
auto pick = [&](std::unique_lock<std::mutex> & lock) -> side * {
|
||||
for (;;)
|
||||
{
|
||||
const auto now = std::chrono::steady_clock::now();
|
||||
side * best = nullptr;
|
||||
int best_score = 0x7fffffff;
|
||||
std::chrono::steady_clock::time_point soonest =
|
||||
now + std::chrono::hours(1);
|
||||
bool any_left = false;
|
||||
for (auto & item : all)
|
||||
{
|
||||
for (int p = 0; p < 2; ++p)
|
||||
{
|
||||
side & s = item.party[p];
|
||||
if (s.cursor.done)
|
||||
continue;
|
||||
if (s.busy)
|
||||
{
|
||||
any_left = true;
|
||||
continue;
|
||||
}
|
||||
any_left = true;
|
||||
if (s.ready_at > now)
|
||||
{
|
||||
if (s.ready_at < soonest)
|
||||
soonest = s.ready_at;
|
||||
continue;
|
||||
}
|
||||
// A side still submitting outranks one parked on a receive.
|
||||
// Among those, the one further behind runs first.
|
||||
const int score = static_cast<int>(s.cursor.exchange_i)
|
||||
+ (s.cursor.awaiting_peer ? 100000 : 0);
|
||||
if (score < best_score)
|
||||
{
|
||||
best_score = score;
|
||||
best = &s;
|
||||
}
|
||||
}
|
||||
}
|
||||
if (best != nullptr)
|
||||
{
|
||||
best->busy = true;
|
||||
return best;
|
||||
}
|
||||
if (!any_left)
|
||||
return nullptr;
|
||||
cv.wait_until(lock, soonest);
|
||||
}
|
||||
};
|
||||
|
||||
auto worker = [&] {
|
||||
std::unique_lock<std::mutex> lock(mu);
|
||||
for (;;)
|
||||
{
|
||||
side * s = pick(lock);
|
||||
if (s == nullptr)
|
||||
return;
|
||||
std::size_t inst_i = 0;
|
||||
int party = 0;
|
||||
for (; inst_i < all.size(); ++inst_i)
|
||||
{
|
||||
if (&all[inst_i].party[0] == s)
|
||||
{
|
||||
party = 0;
|
||||
break;
|
||||
}
|
||||
if (&all[inst_i].party[1] == s)
|
||||
{
|
||||
party = 1;
|
||||
break;
|
||||
}
|
||||
}
|
||||
protocol::drive_options opt;
|
||||
opt.park_if_waiting = true;
|
||||
opt.one_exchange = true;
|
||||
opt.cursor = &s->cursor;
|
||||
auto * sink = s->sink;
|
||||
auto * values = &s->values;
|
||||
const std::size_t step = s->cursor.exchange_i;
|
||||
lock.unlock();
|
||||
|
||||
protocol::drive(plan, *sink, *values, kernels,
|
||||
static_cast<std::size_t>(party), opt);
|
||||
|
||||
lock.lock();
|
||||
s->busy = false;
|
||||
if (s->cursor.done)
|
||||
++finished;
|
||||
else
|
||||
{
|
||||
s->ready_at = std::chrono::steady_clock::now()
|
||||
+ std::chrono::microseconds(chaos_us(inst_i, party, step));
|
||||
}
|
||||
cv.notify_all();
|
||||
}
|
||||
};
|
||||
|
||||
std::vector<std::thread> pool;
|
||||
pool.reserve(workers);
|
||||
for (std::size_t i = 0; i < workers; ++i)
|
||||
pool.emplace_back(worker);
|
||||
for (auto & th : pool)
|
||||
th.join();
|
||||
if (finished != instances * 2)
|
||||
throw std::runtime_error("run_fleet: not every side finished");
|
||||
}
|
||||
|
||||
} // namespace app
|
||||
} // namespace dpf
|
||||
|
||||
#endif
|
||||
312
include/dpf/app_plans.hpp
Normal file
312
include/dpf/app_plans.hpp
Normal file
|
|
@ -0,0 +1,312 @@
|
|||
/// @file dpf/app_plans.hpp
|
||||
/// @brief Named compose plans and star drivers for the application sketches.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_APP_PLANS_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_APP_PLANS_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <functional>
|
||||
#include <memory>
|
||||
#include <stdexcept>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/compose.hpp"
|
||||
#include "dpf/mesh_apps.hpp"
|
||||
#include "dpf/net/edge_mesh.hpp"
|
||||
#include "dpf/pad_graphs.hpp"
|
||||
#include "dpf/protocol.hpp"
|
||||
#include "dpf/session_host.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace protocol
|
||||
{
|
||||
|
||||
/// @brief Drive every session until all instances are done (interleaved).
|
||||
/// @details Callers must `submit` each live index first (see `submit_and_drive`).
|
||||
inline void drive_sessions(std::vector<schedule_session *> sessions,
|
||||
unsigned max_spins = 100000u)
|
||||
{
|
||||
if (sessions.empty())
|
||||
return;
|
||||
for (auto * s : sessions)
|
||||
{
|
||||
if (s == nullptr)
|
||||
throw std::invalid_argument("drive_sessions null");
|
||||
}
|
||||
for (unsigned spins = 0;; ++spins)
|
||||
{
|
||||
if (spins > max_spins)
|
||||
throw std::runtime_error("drive_sessions: peer not ready");
|
||||
bool all_done = true;
|
||||
for (auto * s : sessions)
|
||||
{
|
||||
s->drive();
|
||||
for (std::size_t i = 0; i < s->count(); ++i)
|
||||
all_done = all_done && s->done(i);
|
||||
}
|
||||
if (all_done)
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Submit every unfinished index, then `drive_sessions`.
|
||||
inline void submit_and_drive(std::vector<schedule_session *> sessions,
|
||||
unsigned max_spins = 100000u)
|
||||
{
|
||||
for (auto * s : sessions)
|
||||
{
|
||||
if (s == nullptr)
|
||||
throw std::invalid_argument("submit_and_drive null");
|
||||
for (std::size_t i = 0; i < s->count(); ++i)
|
||||
if (!s->done(i))
|
||||
s->submit(i);
|
||||
}
|
||||
drive_sessions(std::move(sessions), max_spins);
|
||||
}
|
||||
|
||||
/// @brief Client↔N-server star: build sessions from round lists and drive.
|
||||
inline void drive_star(net::memory_star & star,
|
||||
std::vector<schedule_round> client_rounds,
|
||||
const std::function<std::vector<schedule_round>(std::size_t server_i)> &
|
||||
server_rounds_for,
|
||||
unsigned max_spins = 100000u)
|
||||
{
|
||||
schedule_session client(1, star.client_mesh(), std::move(client_rounds),
|
||||
false);
|
||||
std::vector<schedule_session> servers;
|
||||
servers.reserve(star.servers);
|
||||
for (std::size_t i = 0; i < star.servers; ++i)
|
||||
servers.emplace_back(1, star.server_edge(i), server_rounds_for(i),
|
||||
false);
|
||||
std::vector<schedule_session *> ptrs;
|
||||
ptrs.reserve(1 + servers.size());
|
||||
ptrs.push_back(&client);
|
||||
for (auto & s : servers)
|
||||
ptrs.push_back(&s);
|
||||
submit_and_drive(std::move(ptrs), max_spins);
|
||||
}
|
||||
|
||||
/// @brief N-server PIR / keyword PIR compose plan (`client_servers` waves).
|
||||
inline plan n_server_pir_plan(std::size_t party, std::size_t n_servers,
|
||||
std::size_t query_bytes, std::size_t answer_bytes)
|
||||
{
|
||||
composer c(party);
|
||||
auto q = c.client_servers(n_servers, query_bytes, answer_bytes);
|
||||
(void)q;
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
inline plan keyword_pir_compose_plan(std::size_t party, std::size_t depth,
|
||||
std::size_t answer_bytes = sizeof(int))
|
||||
{
|
||||
const std::size_t query_bytes = 16 + depth * 16;
|
||||
return n_server_pir_plan(party, 2, query_bytes, answer_bytes);
|
||||
}
|
||||
|
||||
/// @brief Express / Sabre online walk: fused CW + trailer (sketch/proof).
|
||||
inline plan mailbox_write_fused_plan(std::size_t party, std::size_t depth = 8,
|
||||
std::size_t cw_bytes = 16, std::size_t trailer_bytes = 8)
|
||||
{
|
||||
composer c(party);
|
||||
auto seed = c.input(domain::fss, 16);
|
||||
auto trailer = c.input(domain::a, trailer_bytes);
|
||||
auto walk = c.fss_point_fused(seed, depth, cw_bytes, trailer);
|
||||
(void)walk;
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
/// @brief BitMore: L parallel bit-point walks (CSE → one wave per depth).
|
||||
inline plan bitmore_fan_plan(std::size_t party, std::size_t L = 4,
|
||||
std::size_t depth = 6)
|
||||
{
|
||||
composer c(party);
|
||||
auto seeds = c.fan(L, [&](std::size_t) {
|
||||
return c.input(domain::fss, 16);
|
||||
});
|
||||
(void)c.fan(L, [&](std::size_t i) {
|
||||
return c.fss_point(seeds[i], depth, 16);
|
||||
});
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
/// @brief BitMore / PIRsona star key-ship + answer schedule rounds.
|
||||
inline std::vector<schedule_round> bitmore_star_fetch(std::size_t L,
|
||||
std::size_t seed_bytes, std::size_t answer_bytes,
|
||||
const std::shared_ptr<std::vector<std::vector<std::uint8_t>>> & seeds,
|
||||
const std::shared_ptr<std::vector<std::vector<std::uint8_t>>> & answers)
|
||||
{
|
||||
return pirsona_bitmore_fetch(L, seed_bytes, answer_bytes, seeds, answers);
|
||||
}
|
||||
|
||||
/// @brief SUBLEQ one instruction: prepaid expand + cmp branch skeleton.
|
||||
inline plan subleq_instruction_plan(std::size_t party, std::size_t addr_depth = 8,
|
||||
std::size_t slot_bytes = 16)
|
||||
{
|
||||
composer c(party);
|
||||
auto seed = c.input(domain::fss, 16);
|
||||
auto prepaid = c.defer_expand(seed, addr_depth, slot_bytes);
|
||||
auto branch = c.fss_cmp(seed, addr_depth, slot_bytes);
|
||||
(void)prepaid;
|
||||
(void)branch;
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
/// @brief Queue `n_insns` SUBLEQ instruction plans on a peer `session_host`.
|
||||
inline void subleq_run_instructions(session_host & host, std::size_t n_insns,
|
||||
std::size_t addr_depth = 8, std::size_t slot_bytes = 16)
|
||||
{
|
||||
for (std::size_t i = 0; i < n_insns; ++i)
|
||||
host.push(subleq_instruction_plan(host.party(), addr_depth, slot_bytes));
|
||||
host.drive_until_idle();
|
||||
}
|
||||
|
||||
/// @brief Pika early-stop unit walk (`early_stop` bits pack into the leaf).
|
||||
inline plan pika_lookup_plan(std::size_t party, std::size_t full_depth = 8,
|
||||
std::size_t early_stop = 3)
|
||||
{
|
||||
composer c(party);
|
||||
auto seed = c.input(domain::fss, 16);
|
||||
auto tip = c.fss_point_early_stop(seed, full_depth, early_stop, 16);
|
||||
(void)tip;
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
/// @brief Pika with dealer pad delivery of the unit key before the walk.
|
||||
inline plan pika_dealer_lookup_plan(std::size_t party, std::size_t full_depth = 8,
|
||||
std::size_t early_stop = 3)
|
||||
{
|
||||
composer c(party);
|
||||
auto key = c.input(domain::fss, 16);
|
||||
auto delivered = c.dealer_deliver(key);
|
||||
auto tip = c.fss_point_early_stop(delivered, full_depth, early_stop, 16);
|
||||
(void)tip;
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
/// @brief Duoram write: path CWs with leaf applied later.
|
||||
inline plan duoram_update_plan(std::size_t party, std::size_t depth = 8,
|
||||
std::size_t slot_bytes = 16)
|
||||
{
|
||||
composer c(party);
|
||||
auto seed = c.input(domain::fss, 16);
|
||||
auto walk = c.leaf_later_walk(seed, depth, slot_bytes);
|
||||
auto leaf = c.input(domain::a, slot_bytes);
|
||||
auto applied = c.apply_leaf_correction(walk.values, walk.control, leaf);
|
||||
(void)applied;
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
/// @brief Poplar / Prio / Mastic: prefix checkpoints each level.
|
||||
inline plan poplar_prefix_plan(std::size_t party, std::size_t depth = 8,
|
||||
std::size_t slot_bytes = 16, std::size_t prefix_bytes = 8)
|
||||
{
|
||||
composer c(party);
|
||||
auto seed = c.input(domain::fss, 16);
|
||||
auto prefixes = c.level_walk_prefixes(seed, depth, slot_bytes, prefix_bytes);
|
||||
(void)prefixes;
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
/// @brief Weighted / multi-seed prefix fan (Prio multi-client shape).
|
||||
inline plan poplar_prefix_fan_plan(std::size_t party, std::size_t n_keys,
|
||||
std::size_t depth = 8, std::size_t slot_bytes = 16,
|
||||
std::size_t prefix_bytes = 8)
|
||||
{
|
||||
composer c(party);
|
||||
(void)c.fan(n_keys, [&](std::size_t) {
|
||||
auto seed = c.input(domain::fss, 16);
|
||||
return c.level_walk_prefixes(seed, depth, slot_bytes, prefix_bytes)
|
||||
.at_level.back();
|
||||
});
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
/// @brief Ledger (2,3): verifiable DPF3 append — upload + proof exchange.
|
||||
inline plan ledger23_append_plan(std::size_t party, std::size_t depth = 8,
|
||||
std::size_t proof_bytes = 32)
|
||||
{
|
||||
composer c(party);
|
||||
const std::size_t query_bytes = 16 + depth * 16;
|
||||
auto up = c.client_servers(3, query_bytes, proof_bytes);
|
||||
(void)up;
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
/// @brief Floram DS keygen walk (read/write expand stay local).
|
||||
inline plan floram_ds_plan(std::size_t party, std::size_t depth = 8,
|
||||
std::size_t slot_bytes = 16, bool oh = false)
|
||||
{
|
||||
composer c(party);
|
||||
auto seed = c.input(domain::fss, 16);
|
||||
auto tip = c.level_walk_ds(seed, depth, slot_bytes, oh);
|
||||
(void)tip;
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
/// @brief Splinter-style point walk (server expand cost).
|
||||
inline plan fss_point_plan(std::size_t party, std::size_t depth = 8,
|
||||
std::size_t slot_bytes = 16)
|
||||
{
|
||||
composer c(party);
|
||||
auto seed = c.input(domain::fss, 16);
|
||||
auto tip = c.fss_point(seed, depth, slot_bytes);
|
||||
(void)tip;
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
/// @brief Waldo-style comparison walk.
|
||||
inline plan fss_cmp_plan(std::size_t party, std::size_t depth = 8,
|
||||
std::size_t slot_bytes = 16)
|
||||
{
|
||||
composer c(party);
|
||||
auto seed = c.input(domain::fss, 16);
|
||||
auto tip = c.fss_cmp(seed, depth, slot_bytes);
|
||||
(void)tip;
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
/// @brief Range-count: two parallel cmp walks.
|
||||
inline plan range_count_plan(std::size_t party, std::size_t depth = 8,
|
||||
std::size_t slot_bytes = 16)
|
||||
{
|
||||
composer c(party);
|
||||
auto seed = c.input(domain::fss, 16);
|
||||
(void)c.fan(2, [&](std::size_t) {
|
||||
return c.fss_cmp(seed, depth, slot_bytes);
|
||||
});
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
/// @brief PSI cuckoo probes → multipoint fan (occupied bucket ids).
|
||||
inline plan psi_cuckoo_plan(std::size_t party,
|
||||
const std::vector<std::size_t> & probes, std::size_t depth = 8,
|
||||
std::size_t slot_bytes = 16, std::size_t answer_bytes = 8)
|
||||
{
|
||||
composer c(party);
|
||||
auto mr = c.schedule_cuckoo_probes(probes,
|
||||
[&](std::size_t) { return c.input(domain::fss, 16); }, depth,
|
||||
slot_bytes, answer_bytes);
|
||||
(void)mr;
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
/// @brief idpf_agg: adaptive prefix retain loop (`n_bits` online rounds).
|
||||
inline plan idpf_agg_plan(std::size_t party, std::size_t n_bits = 16,
|
||||
std::size_t slot_bytes = 16, std::size_t prefix_bytes = 8)
|
||||
{
|
||||
composer c(party);
|
||||
auto f = c.begin_adaptive_prefix(c.input(domain::fss, 16));
|
||||
for (std::size_t bit = 0; bit < n_bits; ++bit)
|
||||
{
|
||||
f = c.step_adaptive_prefix(f, slot_bytes, prefix_bytes);
|
||||
f = c.retain_adaptive_prefix(f, 0);
|
||||
}
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
} // namespace protocol
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_APP_PLANS_HPP__
|
||||
144
include/dpf/app_runtime.hpp
Normal file
144
include/dpf/app_runtime.hpp
Normal file
|
|
@ -0,0 +1,144 @@
|
|||
/// @file dpf/app_runtime.hpp
|
||||
/// @brief File prep and two-party synchronous mux helpers.
|
||||
/// @details `stream_runtime` configures only this synchronous TCP mux path.
|
||||
/// For any other transport, party count, or a dealer, use
|
||||
/// `app::run_config` with `app::run_parties` (`party_runner.hpp`).
|
||||
/// The synchronous mux speaks the same frames as the async mux.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_APP_RUNTIME_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_APP_RUNTIME_HPP__
|
||||
|
||||
#include <atomic>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <map>
|
||||
#include <mutex>
|
||||
#include <stdexcept>
|
||||
#include <string>
|
||||
#include <thread>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/compose.hpp"
|
||||
#include "dpf/net/stream_array.hpp"
|
||||
#include "dpf/net/sync_stream_array.hpp"
|
||||
#include "dpf/party_run.hpp"
|
||||
#include "dpf/prep_source.hpp"
|
||||
#include "dpf/protocol_roles.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace app
|
||||
{
|
||||
|
||||
/// @brief File prep basename and synchronous mux settings (two parties).
|
||||
struct stream_runtime
|
||||
{
|
||||
std::string prep_basename;
|
||||
std::string mux_host = "127.0.0.1";
|
||||
unsigned short mux_port = 0;
|
||||
std::size_t mux_nstreams = 1;
|
||||
};
|
||||
|
||||
/// @brief Read one party's prep blob from `basename-p0` / `basename-p1`.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline std::vector<std::uint8_t> open_file_prep(const std::string & basename,
|
||||
unsigned party)
|
||||
{
|
||||
return prep::read_file(basename + "-p" + std::to_string(party));
|
||||
}
|
||||
|
||||
/// @brief Same as `open_file_prep`, parsed as a prep cursor.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline prep::cursor open_file_prep_cursor(const std::string & basename,
|
||||
unsigned party)
|
||||
{
|
||||
return prep::cursor(open_file_prep(basename, party));
|
||||
}
|
||||
|
||||
/// @brief Run `fn(party, mux_stream_array &)` over one TCP connection with mux.
|
||||
template <typename Fn>
|
||||
void with_tcp_mux_peer(unsigned party, std::string host,
|
||||
std::atomic<unsigned short> & port, std::size_t nstreams, Fn && fn)
|
||||
{
|
||||
dpf::run::tcp_pair_mux(party, std::move(host), port, nstreams,
|
||||
std::forward<Fn>(fn));
|
||||
}
|
||||
|
||||
/// @brief Convenience: mux settings from `stream_runtime`.
|
||||
template <typename Fn>
|
||||
void with_tcp_mux_peer(unsigned party, const stream_runtime & rt,
|
||||
std::size_t nstreams, Fn && fn)
|
||||
{
|
||||
std::atomic<unsigned short> port{rt.mux_port};
|
||||
with_tcp_mux_peer(party, rt.mux_host, port, nstreams, std::forward<Fn>(fn));
|
||||
}
|
||||
|
||||
/// @brief One party: TCP mux + `drive_plan_on_streams`.
|
||||
inline void drive_plan_mux(const protocol::plan & plan,
|
||||
net::stream_array & mux, std::vector<std::vector<std::uint8_t>> & values,
|
||||
const std::map<std::uint32_t, protocol::kernel_fn> & kernels, std::size_t party,
|
||||
std::size_t lanes = 1, const protocol::drive_options & opt = {})
|
||||
{
|
||||
protocol::drive_plan_on_streams(plan, mux, values, kernels, party, lanes, opt);
|
||||
}
|
||||
|
||||
/// @brief Two localhost threads on one mux TCP link, both plans driven on streams.
|
||||
inline void drive_plan_mux_both(const protocol::plan & p0,
|
||||
const protocol::plan & p1, std::vector<std::vector<std::uint8_t>> & v0,
|
||||
std::vector<std::vector<std::uint8_t>> & v1,
|
||||
const std::map<std::uint32_t, protocol::kernel_fn> & kernels = {},
|
||||
std::size_t lanes = 1, const protocol::drive_options & opt = {},
|
||||
const stream_runtime & rt = {})
|
||||
{
|
||||
const auto slots0 = p0.slot_bytes_all();
|
||||
const auto slots1 = p1.slot_bytes_all();
|
||||
if (slots0 != slots1)
|
||||
throw std::invalid_argument("drive_plan_mux_both: slot shapes differ");
|
||||
const std::size_t nstreams =
|
||||
slots0.empty() ? 1 : protocol::lanes_for_plan(slots0.size(), opt);
|
||||
std::atomic<unsigned short> port{rt.mux_port};
|
||||
std::exception_ptr err;
|
||||
std::mutex err_mu;
|
||||
auto note = [&](std::exception_ptr e) {
|
||||
std::lock_guard<std::mutex> lock(err_mu);
|
||||
if (!err)
|
||||
err = std::move(e);
|
||||
};
|
||||
std::thread t0([&] {
|
||||
try
|
||||
{
|
||||
with_tcp_mux_peer(0, rt.mux_host, port, nstreams,
|
||||
[&](unsigned, net::mux_stream_array & mux) {
|
||||
drive_plan_mux(p0, mux, v0, kernels, 0, lanes, opt);
|
||||
});
|
||||
}
|
||||
catch (...)
|
||||
{
|
||||
note(std::current_exception());
|
||||
}
|
||||
});
|
||||
std::thread t1([&] {
|
||||
try
|
||||
{
|
||||
with_tcp_mux_peer(1, rt.mux_host, port, nstreams,
|
||||
[&](unsigned, net::mux_stream_array & mux) {
|
||||
drive_plan_mux(p1, mux, v1, kernels, 1, lanes, opt);
|
||||
});
|
||||
}
|
||||
catch (...)
|
||||
{
|
||||
note(std::current_exception());
|
||||
}
|
||||
});
|
||||
t0.join();
|
||||
t1.join();
|
||||
if (err)
|
||||
std::rethrow_exception(err);
|
||||
}
|
||||
|
||||
} // namespace app
|
||||
} // namespace dpf
|
||||
|
||||
#endif
|
||||
777
include/dpf/arith_garble.hpp
Normal file
777
include/dpf/arith_garble.hpp
Normal file
|
|
@ -0,0 +1,777 @@
|
|||
/// @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__
|
||||
|
|
@ -68,6 +68,10 @@ auto async_post(ExecutorT executor, Function && func, CompletionToken && token)
|
|||
// make_dpf
|
||||
//
|
||||
|
||||
/// \complexity Local `make_dpf` is O(n) per key, then one write of each key. n is `depth`.
|
||||
/// \rounds 1 per key. Each party is a single `asio::write` of six buffers; there is no reply.
|
||||
/// \communication Per key, two copies (one per peer) of the correction-word array (n nodes), the advice array (n bytes), one root, the leaf tuple, the beaver tuple, and the offset word.
|
||||
/// \preprocessing none beyond the local keygen. The dealer holds the clear point.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename PeerT,
|
||||
|
|
@ -123,6 +127,10 @@ auto make_dpf(PeerT & peer0, PeerT & peer1, std::size_t count, dpfargs<InputT, O
|
|||
return std::make_tuple(bytes_written0, bytes_written1, count);
|
||||
}
|
||||
|
||||
/// \complexity Local `make_dpf` is O(n) per key, then one write of each key. n is `depth`.
|
||||
/// \rounds 1 per key. Each party is a single `asio::write` of six buffers; there is no reply.
|
||||
/// \communication Per key, two copies (one per peer) of the correction-word array (n nodes), the advice array (n bytes), one root, the leaf tuple, the beaver tuple, and the offset word.
|
||||
/// \preprocessing none beyond the local keygen. The dealer holds the clear point.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename PeerT,
|
||||
|
|
@ -138,6 +146,10 @@ auto make_dpf(PeerT & peer0, PeerT & peer1, std::size_t count, dpfargs<InputT, O
|
|||
return ret;
|
||||
}
|
||||
|
||||
/// \complexity Local `make_dpf` is O(n) per key, then one write of each key. n is `depth`.
|
||||
/// \rounds 1 per key. Each party is a single `asio::write` of six buffers; there is no reply.
|
||||
/// \communication Per key, two copies (one per peer) of the correction-word array (n nodes), the advice array (n bytes), one root, the leaf tuple, the beaver tuple, and the offset word.
|
||||
/// \preprocessing none beyond the local keygen. The dealer holds the clear point.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename PeerT,
|
||||
|
|
@ -152,6 +164,10 @@ auto make_dpf(PeerT & peer0, PeerT & peer1, dpfargs<InputT, OutputT, OutputTs...
|
|||
return std::make_tuple(bytes_written0, bytes_written1);
|
||||
}
|
||||
|
||||
/// \complexity Local `make_dpf` is O(n) per key, then one write of each key. n is `depth`.
|
||||
/// \rounds 1 per key. Each party is a single `asio::write` of six buffers; there is no reply.
|
||||
/// \communication Per key, two copies (one per peer) of the correction-word array (n nodes), the advice array (n bytes), one root, the leaf tuple, the beaver tuple, and the offset word.
|
||||
/// \preprocessing none beyond the local keygen. The dealer holds the clear point.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename PeerT,
|
||||
|
|
@ -646,6 +662,10 @@ auto async_read_dpf(DealerT & dealer, Emplaceable & output, CompletionToken && t
|
|||
// assign_wildcard_input
|
||||
//
|
||||
|
||||
/// \complexity O(1) arithmetic besides the socket transfer.
|
||||
/// \rounds 1. One write of the local share, one read of the peer share (`async_assign_wildcard_input`).
|
||||
/// \communication `sizeof(input_type)` bytes each way.
|
||||
/// \preprocessing The mask in `offset_x` was sampled at `make_dpf`. This exchange opens mask − alpha.
|
||||
template <typename PeerT,
|
||||
typename DpfKey,
|
||||
typename InputType>
|
||||
|
|
@ -675,6 +695,10 @@ auto assign_wildcard_input(PeerT & peer_in, PeerT & peer_out, DpfKey & dpf,
|
|||
return std::make_tuple(offset_share, bytes_written, bytes_read);
|
||||
}
|
||||
|
||||
/// \complexity O(1) arithmetic besides the socket transfer.
|
||||
/// \rounds 1. One write of the local share, one read of the peer share (`async_assign_wildcard_input`).
|
||||
/// \communication `sizeof(input_type)` bytes each way.
|
||||
/// \preprocessing The mask in `offset_x` was sampled at `make_dpf`. This exchange opens mask − alpha.
|
||||
template <typename PeerT,
|
||||
typename DpfKey,
|
||||
typename InputType>
|
||||
|
|
@ -686,6 +710,10 @@ auto assign_wildcard_input(PeerT & peer, DpfKey & dpf, InputType && input_share,
|
|||
std::forward<InputType>(input_share), error);
|
||||
}
|
||||
|
||||
/// \complexity O(1) arithmetic besides the socket transfer.
|
||||
/// \rounds 1. One write of the local share, one read of the peer share (`async_assign_wildcard_input`).
|
||||
/// \communication `sizeof(input_type)` bytes each way.
|
||||
/// \preprocessing The mask in `offset_x` was sampled at `make_dpf`. This exchange opens mask − alpha.
|
||||
template <typename PeerT,
|
||||
typename DpfKey,
|
||||
typename InputType>
|
||||
|
|
@ -700,6 +728,10 @@ auto assign_wildcard_input(PeerT & peer_in, PeerT & peer_out, DpfKey & dpf,
|
|||
return ret;
|
||||
}
|
||||
|
||||
/// \complexity O(1) arithmetic besides the socket transfer.
|
||||
/// \rounds 1. One write of the local share, one read of the peer share (`async_assign_wildcard_input`).
|
||||
/// \communication `sizeof(input_type)` bytes each way.
|
||||
/// \preprocessing The mask in `offset_x` was sampled at `make_dpf`. This exchange opens mask − alpha.
|
||||
template <typename PeerT,
|
||||
typename DpfKey,
|
||||
typename InputType>
|
||||
|
|
@ -717,6 +749,10 @@ auto assign_wildcard_input(PeerT & peer, DpfKey & dpf, InputType && input_share)
|
|||
// async_assign_wildcard_input
|
||||
//
|
||||
|
||||
/// \complexity O(1) arithmetic besides the socket transfer.
|
||||
/// \rounds 1. One write of the local share, one read of the peer share (`async_assign_wildcard_input`).
|
||||
/// \communication `sizeof(input_type)` bytes each way.
|
||||
/// \preprocessing The mask in `offset_x` was sampled at `make_dpf`. This exchange opens mask − alpha.
|
||||
template <typename PeerT,
|
||||
typename ExecutorT,
|
||||
typename DpfKey,
|
||||
|
|
@ -796,6 +832,10 @@ auto async_assign_wildcard_input(PeerT & peer_in, PeerT & peer_out,
|
|||
#include <asio/unyield.hpp>
|
||||
}
|
||||
|
||||
/// \complexity O(1) arithmetic besides the socket transfer.
|
||||
/// \rounds 1. One write of the local share, one read of the peer share (`async_assign_wildcard_input`).
|
||||
/// \communication `sizeof(input_type)` bytes each way.
|
||||
/// \preprocessing The mask in `offset_x` was sampled at `make_dpf`. This exchange opens mask − alpha.
|
||||
template <typename PeerT,
|
||||
typename ExecutorT,
|
||||
typename DpfKey,
|
||||
|
|
@ -811,6 +851,10 @@ auto async_assign_wildcard_input(PeerT & peer, ExecutorT work_executor,
|
|||
std::forward<CompletionToken>(token));
|
||||
}
|
||||
|
||||
/// \complexity O(1) arithmetic besides the socket transfer.
|
||||
/// \rounds 1. One write of the local share, one read of the peer share (`async_assign_wildcard_input`).
|
||||
/// \communication `sizeof(input_type)` bytes each way.
|
||||
/// \preprocessing The mask in `offset_x` was sampled at `make_dpf`. This exchange opens mask − alpha.
|
||||
template <typename PeerT,
|
||||
typename DpfKey,
|
||||
typename InputType,
|
||||
|
|
@ -826,6 +870,10 @@ auto async_assign_wildcard_input(PeerT & peer_in, PeerT & peer_out,
|
|||
std::forward<CompletionToken>(token));
|
||||
}
|
||||
|
||||
/// \complexity O(1) arithmetic besides the socket transfer.
|
||||
/// \rounds 1. One write of the local share, one read of the peer share (`async_assign_wildcard_input`).
|
||||
/// \communication `sizeof(input_type)` bytes each way.
|
||||
/// \preprocessing The mask in `offset_x` was sampled at `make_dpf`. This exchange opens mask − alpha.
|
||||
template <typename PeerT,
|
||||
typename DpfKey,
|
||||
typename InputType,
|
||||
|
|
@ -844,6 +892,10 @@ auto async_assign_wildcard_input(PeerT & peer, DpfKey & dpf, InputType && input_
|
|||
// assign_wildcard_output
|
||||
//
|
||||
|
||||
/// \complexity O(leaf bytes) for the Beaver leaf arithmetic, plus the transfers.
|
||||
/// \rounds 2. Write/read the blinded output share, then write/read the leaf share (`async_assign_wildcard_output`).
|
||||
/// \communication `sizeof(output_type)` plus `sizeof(leaf_type)` each way.
|
||||
/// \preprocessing The leaf Beaver triple (`vector_blind`, `output_blind`, `blinded_vector`) was stored at keygen.
|
||||
template <std::size_t I = 0,
|
||||
typename PeerT,
|
||||
typename DpfKey,
|
||||
|
|
@ -861,7 +913,11 @@ auto assign_wildcard_output(PeerT & peer_in, PeerT & peer_out, DpfKey & dpf,
|
|||
constexpr bool is_packed = true;
|
||||
|
||||
auto & leaf_wrapper = utils::get<I>(dpf.leaf_nodes);
|
||||
|
||||
|
||||
// Second (and later) assigns install β'−β on top of the ready leaf.
|
||||
if (leaf_wrapper.is_ready())
|
||||
leaf_wrapper.begin_update();
|
||||
|
||||
auto blinded_output = leaf_wrapper.compute_and_get_blinded_output_share(output_share);
|
||||
bytes_written += ::asio::write(peer_out, ::asio::buffer(&blinded_output, sizeof(output_type)), error);
|
||||
if (error)
|
||||
|
|
@ -898,6 +954,10 @@ auto assign_wildcard_output(PeerT & peer_in, PeerT & peer_out, DpfKey & dpf,
|
|||
return std::make_tuple(leaf_share, bytes_written, bytes_read);
|
||||
}
|
||||
|
||||
/// \complexity O(leaf bytes) for the Beaver leaf arithmetic, plus the transfers.
|
||||
/// \rounds 2. Write/read the blinded output share, then write/read the leaf share (`async_assign_wildcard_output`).
|
||||
/// \communication `sizeof(output_type)` plus `sizeof(leaf_type)` each way.
|
||||
/// \preprocessing The leaf Beaver triple (`vector_blind`, `output_blind`, `blinded_vector`) was stored at keygen.
|
||||
template <std::size_t I = 0,
|
||||
typename PeerT,
|
||||
typename DpfKey,
|
||||
|
|
@ -910,6 +970,10 @@ auto assign_wildcard_output(PeerT & peer, DpfKey & dpf,
|
|||
std::forward<OutputType>(output_share), error);
|
||||
}
|
||||
|
||||
/// \complexity O(leaf bytes) for the Beaver leaf arithmetic, plus the transfers.
|
||||
/// \rounds 2. Write/read the blinded output share, then write/read the leaf share (`async_assign_wildcard_output`).
|
||||
/// \communication `sizeof(output_type)` plus `sizeof(leaf_type)` each way.
|
||||
/// \preprocessing The leaf Beaver triple (`vector_blind`, `output_blind`, `blinded_vector`) was stored at keygen.
|
||||
template <std::size_t I = 0,
|
||||
typename PeerT,
|
||||
typename DpfKey,
|
||||
|
|
@ -925,6 +989,10 @@ auto assign_wildcard_output(PeerT & peer_in, PeerT & peer_out, DpfKey & dpf,
|
|||
return ret;
|
||||
}
|
||||
|
||||
/// \complexity O(leaf bytes) for the Beaver leaf arithmetic, plus the transfers.
|
||||
/// \rounds 2. Write/read the blinded output share, then write/read the leaf share (`async_assign_wildcard_output`).
|
||||
/// \communication `sizeof(output_type)` plus `sizeof(leaf_type)` each way.
|
||||
/// \preprocessing The leaf Beaver triple (`vector_blind`, `output_blind`, `blinded_vector`) was stored at keygen.
|
||||
template <std::size_t I = 0,
|
||||
typename PeerT,
|
||||
typename DpfKey,
|
||||
|
|
@ -943,6 +1011,10 @@ auto assign_wildcard_output(PeerT & peer, DpfKey & dpf, OutputType && output_sha
|
|||
// async_assign_wildcard_output
|
||||
//
|
||||
|
||||
/// \complexity O(leaf bytes) for the Beaver leaf arithmetic, plus the transfers.
|
||||
/// \rounds 2. Write/read the blinded output share, then write/read the leaf share (`async_assign_wildcard_output`).
|
||||
/// \communication `sizeof(output_type)` plus `sizeof(leaf_type)` each way.
|
||||
/// \preprocessing The leaf Beaver triple (`vector_blind`, `output_blind`, `blinded_vector`) was stored at keygen.
|
||||
template <std::size_t I = 0,
|
||||
typename PeerT,
|
||||
typename ExecutorT,
|
||||
|
|
@ -987,6 +1059,8 @@ auto async_assign_wildcard_output(PeerT & peer_in, PeerT & peer_out,
|
|||
{
|
||||
yield async_post(work_executor, [&leaf, output_share]() mutable
|
||||
{
|
||||
if (leaf.is_ready())
|
||||
leaf.begin_update();
|
||||
*output_share = leaf.compute_and_get_blinded_output_share(*output_share);
|
||||
}, std::move(self));
|
||||
|
||||
|
|
@ -1059,6 +1133,10 @@ auto async_assign_wildcard_output(PeerT & peer_in, PeerT & peer_out,
|
|||
#include <asio/unyield.hpp>
|
||||
}
|
||||
|
||||
/// \complexity O(leaf bytes) for the Beaver leaf arithmetic, plus the transfers.
|
||||
/// \rounds 2. Write/read the blinded output share, then write/read the leaf share (`async_assign_wildcard_output`).
|
||||
/// \communication `sizeof(output_type)` plus `sizeof(leaf_type)` each way.
|
||||
/// \preprocessing The leaf Beaver triple (`vector_blind`, `output_blind`, `blinded_vector`) was stored at keygen.
|
||||
template <std::size_t I = 0,
|
||||
typename PeerT,
|
||||
typename ExecutorT,
|
||||
|
|
@ -1075,6 +1153,10 @@ auto async_assign_wildcard_output(PeerT & peer, ExecutorT work_executor,
|
|||
std::forward<CompletionToken>(token));
|
||||
}
|
||||
|
||||
/// \complexity O(leaf bytes) for the Beaver leaf arithmetic, plus the transfers.
|
||||
/// \rounds 2. Write/read the blinded output share, then write/read the leaf share (`async_assign_wildcard_output`).
|
||||
/// \communication `sizeof(output_type)` plus `sizeof(leaf_type)` each way.
|
||||
/// \preprocessing The leaf Beaver triple (`vector_blind`, `output_blind`, `blinded_vector`) was stored at keygen.
|
||||
template <std::size_t I = 0,
|
||||
typename PeerT,
|
||||
typename DpfKey,
|
||||
|
|
@ -1091,6 +1173,10 @@ auto async_assign_wildcard_output(PeerT & peer_in, PeerT & peer_out,
|
|||
std::forward<CompletionToken>(token));
|
||||
}
|
||||
|
||||
/// \complexity O(leaf bytes) for the Beaver leaf arithmetic, plus the transfers.
|
||||
/// \rounds 2. Write/read the blinded output share, then write/read the leaf share (`async_assign_wildcard_output`).
|
||||
/// \communication `sizeof(output_type)` plus `sizeof(leaf_type)` each way.
|
||||
/// \preprocessing The leaf Beaver triple (`vector_blind`, `output_blind`, `blinded_vector`) was stored at keygen.
|
||||
template <std::size_t I = 0,
|
||||
typename PeerT,
|
||||
typename DpfKey,
|
||||
|
|
@ -1101,7 +1187,6 @@ HEDLEY_ALWAYS_INLINE
|
|||
auto async_assign_wildcard_output(PeerT & peer, DpfKey & dpf,
|
||||
OutputType && output_share, CompletionToken && token)
|
||||
{
|
||||
auto work_executor = ::asio::system_executor();
|
||||
return async_assign_wildcard_output<I>(peer, peer, dpf,
|
||||
std::forward<OutputType>(output_share),
|
||||
std::forward<CompletionToken>(token));
|
||||
|
|
|
|||
234
include/dpf/async_protocol.hpp
Normal file
234
include/dpf/async_protocol.hpp
Normal file
|
|
@ -0,0 +1,234 @@
|
|||
/// @file dpf/async_protocol.hpp
|
||||
/// @brief Real overlapped, event-driven byte-round protocol runner.
|
||||
/// @details `overlapped_byte_protocol` replaces the blocking loop of
|
||||
/// `factory::async_byte_protocol` with true asynchronous I/O over an
|
||||
/// `async_stream_array`. Each round: read the blind from the dealer
|
||||
/// (async, skipped when `blind_bytes == 0`), `produce` the outbound
|
||||
/// message, then overlap the peer write and peer read as one
|
||||
/// `async_exchange`, and on completion `finish`, fire the round
|
||||
/// callback, and continue to the next round or the done handler.
|
||||
/// Nothing spins on a `peer_ready` flag and nothing blocks the calling
|
||||
/// thread — every continuation runs on an `io_context` thread.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_ASYNC_PROTOCOL_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_ASYNC_PROTOCOL_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <functional>
|
||||
#include <memory>
|
||||
#include <mutex>
|
||||
#include <stdexcept>
|
||||
#include <system_error>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/net/asio_ns.hpp"
|
||||
|
||||
#include "dpf/net/async_stream_array.hpp"
|
||||
#include "dpf/protocol_factory.hpp" // factory::async_byte_round
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace async
|
||||
{
|
||||
|
||||
/// @brief Overlap a write and a read on stream `i`; fire `done` once both end.
|
||||
/// @details The two operations are issued concurrently (full duplex): the
|
||||
/// handler runs after both complete, carrying the first error seen.
|
||||
/// `out` and `in` must stay valid until `done` fires.
|
||||
inline void async_exchange(net::async_stream_array & s, std::size_t i,
|
||||
const void * out, std::size_t out_n, void * in, std::size_t in_n,
|
||||
net::async_handler done)
|
||||
{
|
||||
struct state
|
||||
{
|
||||
net::async_handler done;
|
||||
int remaining = 2;
|
||||
std::error_code ec;
|
||||
std::mutex mu;
|
||||
};
|
||||
auto st = std::make_shared<state>();
|
||||
st->done = std::move(done);
|
||||
auto complete = [st](const std::error_code & ec) {
|
||||
bool fire = false;
|
||||
std::error_code final_ec;
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(st->mu);
|
||||
if (ec && !st->ec)
|
||||
st->ec = ec;
|
||||
if (--st->remaining == 0)
|
||||
fire = true;
|
||||
final_ec = st->ec;
|
||||
}
|
||||
if (fire && st->done)
|
||||
st->done(final_ec);
|
||||
};
|
||||
s.async_write(i, out, out_n, complete);
|
||||
s.async_read(i, in, in_n, complete);
|
||||
}
|
||||
|
||||
/// @brief Fully overlapped runner for a list of `async_byte_round`s.
|
||||
/// @details Constructed per party per run and driven through a `shared_ptr`
|
||||
/// (kept alive by its own continuations). `start` kicks round 0.
|
||||
class overlapped_byte_protocol
|
||||
: public std::enable_shared_from_this<overlapped_byte_protocol>
|
||||
{
|
||||
public:
|
||||
/// @param ec error (empty on success); @param state final protocol state.
|
||||
using done_handler =
|
||||
std::function<void(const std::error_code & ec,
|
||||
std::vector<std::uint8_t> state)>;
|
||||
using on_round_complete_fn = std::function<void(std::size_t round)>;
|
||||
|
||||
/// @param peer overlapped peer array; lane = `r % peer.size()` so a small
|
||||
/// pool can carry many sequential rounds (known-size exchange).
|
||||
/// @param dealer optional blind source (lane `r % dealer.size()`); may be
|
||||
/// null when every round has `blind_bytes == 0`.
|
||||
overlapped_byte_protocol(net::async_stream_array & peer,
|
||||
net::async_stream_array * dealer,
|
||||
std::vector<factory::async_byte_round> rounds,
|
||||
on_round_complete_fn on_round_complete = {})
|
||||
: peer_(&peer),
|
||||
dealer_(dealer),
|
||||
rounds_(std::move(rounds)),
|
||||
on_round_complete_(std::move(on_round_complete))
|
||||
{
|
||||
if (peer_->size() == 0)
|
||||
throw std::invalid_argument("overlapped_byte_protocol: empty peer");
|
||||
for (const auto & r : rounds_)
|
||||
{
|
||||
if (!r.produce || !r.finish)
|
||||
throw std::invalid_argument("overlapped_byte_protocol: missing callbacks");
|
||||
if (r.blind_bytes != 0 && dealer_ == nullptr)
|
||||
throw std::invalid_argument("overlapped_byte_protocol: blind needs a dealer");
|
||||
}
|
||||
if (dealer_ != nullptr && dealer_->size() == 0)
|
||||
throw std::invalid_argument("overlapped_byte_protocol: empty dealer");
|
||||
}
|
||||
|
||||
std::size_t rounds() const noexcept { return rounds_.size(); }
|
||||
|
||||
/// @brief Begin the protocol. `done` fires once, after the last round.
|
||||
void start(std::size_t index, std::vector<std::uint8_t> state,
|
||||
done_handler done)
|
||||
{
|
||||
(void)index;
|
||||
state_ = std::move(state);
|
||||
done_ = std::move(done);
|
||||
run_round(0);
|
||||
}
|
||||
|
||||
private:
|
||||
void finish_with(const std::error_code & ec)
|
||||
{
|
||||
if (done_)
|
||||
{
|
||||
auto d = std::move(done_);
|
||||
done_ = nullptr;
|
||||
d(ec, std::move(state_));
|
||||
}
|
||||
}
|
||||
|
||||
void run_round(std::size_t r)
|
||||
{
|
||||
if (r >= rounds_.size())
|
||||
{
|
||||
finish_with(std::error_code{});
|
||||
return;
|
||||
}
|
||||
const auto & spec = rounds_[r];
|
||||
blind_.assign(spec.blind_bytes, 0);
|
||||
auto self = shared_from_this();
|
||||
if (spec.blind_bytes != 0)
|
||||
{
|
||||
const std::size_t dlane = r % dealer_->size();
|
||||
dealer_->async_read(dlane, blind_.data(), blind_.size(),
|
||||
[self, r](const std::error_code & ec) {
|
||||
self->after_blind(r, ec);
|
||||
});
|
||||
}
|
||||
else
|
||||
{
|
||||
// No blind: hop through the io_context so we never recurse deeply.
|
||||
asio::post(peer_->context(),
|
||||
[self, r]() { self->after_blind(r, std::error_code{}); });
|
||||
}
|
||||
}
|
||||
|
||||
void after_blind(std::size_t r, const std::error_code & ec)
|
||||
{
|
||||
if (ec)
|
||||
{
|
||||
finish_with(ec);
|
||||
return;
|
||||
}
|
||||
const auto & spec = rounds_[r];
|
||||
outbound_ = spec.produce(state_, blind_.data(), blind_.size());
|
||||
if (outbound_.size() != spec.msg_bytes)
|
||||
{
|
||||
finish_with(std::make_error_code(std::errc::message_size));
|
||||
return;
|
||||
}
|
||||
inbound_.assign(spec.msg_bytes, 0);
|
||||
auto self = shared_from_this();
|
||||
const std::size_t lane = r % peer_->size();
|
||||
async_exchange(*peer_, lane, outbound_.data(), outbound_.size(),
|
||||
inbound_.data(), inbound_.size(),
|
||||
[self, r](const std::error_code & xec) {
|
||||
self->after_exchange(r, xec);
|
||||
});
|
||||
}
|
||||
|
||||
void after_exchange(std::size_t r, const std::error_code & ec)
|
||||
{
|
||||
if (ec)
|
||||
{
|
||||
finish_with(ec);
|
||||
return;
|
||||
}
|
||||
const auto & spec = rounds_[r];
|
||||
spec.finish(state_, inbound_.data(), inbound_.size(), blind_.data(),
|
||||
blind_.size());
|
||||
if (on_round_complete_)
|
||||
on_round_complete_(r);
|
||||
run_round(r + 1);
|
||||
}
|
||||
|
||||
net::async_stream_array * peer_ = nullptr;
|
||||
net::async_stream_array * dealer_ = nullptr;
|
||||
std::vector<factory::async_byte_round> rounds_;
|
||||
on_round_complete_fn on_round_complete_;
|
||||
|
||||
std::vector<std::uint8_t> state_;
|
||||
done_handler done_;
|
||||
std::vector<std::uint8_t> blind_;
|
||||
std::vector<std::uint8_t> outbound_;
|
||||
std::vector<std::uint8_t> inbound_;
|
||||
};
|
||||
|
||||
/// @brief Make an `overlapped_byte_protocol` as a `shared_ptr` (required for
|
||||
/// the self-owning continuation chain).
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline std::shared_ptr<overlapped_byte_protocol> make_overlapped_byte_protocol(
|
||||
net::async_stream_array & peer, net::async_stream_array * dealer,
|
||||
std::vector<factory::async_byte_round> rounds,
|
||||
overlapped_byte_protocol::on_round_complete_fn on_round_complete = {})
|
||||
{
|
||||
return std::make_shared<overlapped_byte_protocol>(peer, dealer,
|
||||
std::move(rounds), std::move(on_round_complete));
|
||||
}
|
||||
|
||||
/// @brief Post `start_all`, then run the io_context until all work drains.
|
||||
/// @details The single entry point for driving overlapped parties: everything
|
||||
/// the parties initiate is chased to completion by `io.run()`.
|
||||
template <typename Fn>
|
||||
void run_overlapped(asio::io_context & io, Fn start_all)
|
||||
{
|
||||
asio::post(io, [start_all = std::move(start_all)]() mutable { start_all(); });
|
||||
io.run();
|
||||
}
|
||||
|
||||
} // namespace async
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_ASYNC_PROTOCOL_HPP__
|
||||
File diff suppressed because it is too large
Load diff
578
include/dpf/bench_cells.hpp
Normal file
578
include/dpf/bench_cells.hpp
Normal file
|
|
@ -0,0 +1,578 @@
|
|||
/// @file dpf/bench_cells.hpp
|
||||
/// @brief Bench cells for work that is not a DPF walk.
|
||||
/// @details A secret index stays a key. These cells time what happens after
|
||||
/// the parties already hold shares: a constant-round word gadget, a
|
||||
/// stacked branch of a leaf netlist, a short public table, and a
|
||||
/// hidden reorder of an RSS column. Each cell is a compose plan. The
|
||||
/// harness drives it with `run_parties`, so the payload crosses the
|
||||
/// same transport as a DPF plan (`DPF_TRANSPORT`).
|
||||
///
|
||||
/// Two-party cells put the garble or the table on party 0 and the
|
||||
/// second evaluation on party 1, with one peer payload between them.
|
||||
/// The shuffle is three ring passes. The left-out party's slot is
|
||||
/// empty; the other two carry that pass's array.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_BENCH_CELLS_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_BENCH_CELLS_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <numeric>
|
||||
#include <stdexcept>
|
||||
#include <string>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/arith_garble.hpp"
|
||||
#include "dpf/compose.hpp"
|
||||
#include "dpf/flute.hpp"
|
||||
#include "dpf/shuffle.hpp"
|
||||
#include "dpf/yao.hpp"
|
||||
#include "dpf/yao_stack.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace bench
|
||||
{
|
||||
|
||||
/// @brief One measured configuration. The low 16 bits ride in the node aux.
|
||||
enum class kind : std::uint32_t
|
||||
{
|
||||
proj5 = 1,
|
||||
proj17,
|
||||
proj64,
|
||||
mul5,
|
||||
mul7,
|
||||
mul11,
|
||||
thresh8,
|
||||
thresh16,
|
||||
chain7,
|
||||
yao_if4,
|
||||
yao_if16,
|
||||
yao_hot4,
|
||||
yao_hot8,
|
||||
flute2,
|
||||
flute4,
|
||||
flute8,
|
||||
flute4x8,
|
||||
shuf16,
|
||||
shuf64,
|
||||
shuf256
|
||||
};
|
||||
|
||||
inline constexpr std::uint32_t aux_pass(kind k, unsigned pass) noexcept
|
||||
{
|
||||
return static_cast<std::uint32_t>(k) | (static_cast<std::uint32_t>(pass) << 16);
|
||||
}
|
||||
|
||||
inline kind kind_of(std::uint32_t aux) noexcept
|
||||
{
|
||||
return static_cast<kind>(aux & 0xffffu);
|
||||
}
|
||||
|
||||
inline unsigned pass_of(std::uint32_t aux) noexcept
|
||||
{
|
||||
return aux >> 16;
|
||||
}
|
||||
|
||||
inline bool is_shuffle(kind k) noexcept
|
||||
{
|
||||
return k == kind::shuf16 || k == kind::shuf64 || k == kind::shuf256;
|
||||
}
|
||||
|
||||
inline std::size_t shuffle_n(kind k)
|
||||
{
|
||||
if (k == kind::shuf16)
|
||||
return 16;
|
||||
if (k == kind::shuf64)
|
||||
return 64;
|
||||
if (k == kind::shuf256)
|
||||
return 256;
|
||||
throw std::invalid_argument("bench: shuffle cell");
|
||||
}
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
inline dpf::yao::netlist and_n(unsigned n)
|
||||
{
|
||||
dpf::yao::netlist nl;
|
||||
std::vector<dpf::yao::bit> in;
|
||||
for (unsigned i = 0; i < n; ++i)
|
||||
in.push_back(nl.shared_in());
|
||||
auto acc = in[0];
|
||||
for (unsigned i = 1; i < n; ++i)
|
||||
acc = nl.and_(acc, in[i]);
|
||||
nl.out(acc);
|
||||
return nl;
|
||||
}
|
||||
|
||||
inline dpf::yao::netlist xor2()
|
||||
{
|
||||
dpf::yao::netlist nl;
|
||||
auto a = nl.shared_in();
|
||||
auto b = nl.shared_in();
|
||||
nl.out(nl.xor_(a, b));
|
||||
return nl;
|
||||
}
|
||||
|
||||
inline std::uint64_t mix_bytes(const std::uint8_t * p, std::size_t n)
|
||||
{
|
||||
std::uint64_t h = 14695981039346656037ull;
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
h ^= p[i];
|
||||
h *= 1099511628211ull;
|
||||
}
|
||||
return h;
|
||||
}
|
||||
|
||||
inline void paint(std::uint8_t * out, std::size_t n, std::uint64_t mix)
|
||||
{
|
||||
for (std::size_t i = 0; i < n; i += 8)
|
||||
{
|
||||
const std::uint64_t w = mix + static_cast<std::uint64_t>(i);
|
||||
const std::size_t k = std::min<std::size_t>(8, n - i);
|
||||
std::memcpy(out + i, &w, k);
|
||||
}
|
||||
}
|
||||
|
||||
inline std::size_t proj_bytes(std::uint16_t mod)
|
||||
{
|
||||
arith_garble::circuit c;
|
||||
auto x = c.input(mod);
|
||||
std::vector<std::uint16_t> id(mod);
|
||||
for (std::uint16_t i = 0; i < mod; ++i)
|
||||
id[i] = i;
|
||||
c.out(c.project(x, mod, std::move(id)));
|
||||
const std::uint16_t in = 1;
|
||||
return std::max<std::size_t>(
|
||||
32, dpf::arith_garble::eval_pair(c, &in).ciphertext_rows * 32u);
|
||||
}
|
||||
|
||||
inline std::size_t mul_bytes(std::uint16_t p)
|
||||
{
|
||||
arith_garble::circuit c;
|
||||
auto x = c.input(p);
|
||||
auto y = c.input(p);
|
||||
c.out(c.mul(x, y));
|
||||
const std::uint16_t in[2] = {1, 1};
|
||||
return std::max<std::size_t>(
|
||||
32, dpf::arith_garble::eval_pair(c, in).ciphertext_rows * 32u);
|
||||
}
|
||||
|
||||
inline std::size_t thresh_bytes(unsigned b)
|
||||
{
|
||||
arith_garble::circuit c;
|
||||
std::vector<arith_garble::wire> bits;
|
||||
const auto mod = static_cast<std::uint16_t>(b + 1);
|
||||
for (unsigned i = 0; i < b; ++i)
|
||||
bits.push_back(c.input(mod));
|
||||
c.out(c.threshold(bits, static_cast<std::uint16_t>(b)));
|
||||
std::vector<std::uint16_t> in(b, 1);
|
||||
return std::max<std::size_t>(
|
||||
32, dpf::arith_garble::eval_pair(c, in.data()).ciphertext_rows * 32u);
|
||||
}
|
||||
|
||||
inline std::size_t chain_bytes()
|
||||
{
|
||||
arith_garble::circuit c;
|
||||
auto x = c.input(7);
|
||||
auto y = c.input(7);
|
||||
auto z = c.mul(x, y);
|
||||
for (int i = 0; i < 3; ++i)
|
||||
z = c.mul(z, x);
|
||||
c.out(z);
|
||||
const std::uint16_t in[2] = {2, 3};
|
||||
return std::max<std::size_t>(
|
||||
32, dpf::arith_garble::eval_pair(c, in).ciphertext_rows * 32u);
|
||||
}
|
||||
|
||||
inline std::uint64_t run_proj(std::uint16_t mod)
|
||||
{
|
||||
arith_garble::circuit c;
|
||||
auto x = c.input(mod);
|
||||
std::vector<std::uint16_t> id(mod);
|
||||
for (std::uint16_t i = 0; i < mod; ++i)
|
||||
id[i] = static_cast<std::uint16_t>((i * 3) % mod);
|
||||
c.out(c.project(x, mod, std::move(id)));
|
||||
const std::uint16_t in = static_cast<std::uint16_t>(mod / 2);
|
||||
auto got = dpf::arith_garble::eval_pair(c, &in);
|
||||
return got.opened.empty() ? 0 : got.opened[0];
|
||||
}
|
||||
|
||||
inline std::uint64_t run_mul(std::uint16_t p)
|
||||
{
|
||||
arith_garble::circuit c;
|
||||
auto x = c.input(p);
|
||||
auto y = c.input(p);
|
||||
c.out(c.mul(x, y));
|
||||
const std::uint16_t in[2] = {static_cast<std::uint16_t>(p - 1), 2};
|
||||
auto got = dpf::arith_garble::eval_pair(c, in);
|
||||
return got.opened.empty() ? 0 : got.opened[0];
|
||||
}
|
||||
|
||||
inline std::uint64_t run_thresh(unsigned b)
|
||||
{
|
||||
arith_garble::circuit c;
|
||||
std::vector<arith_garble::wire> bits;
|
||||
const auto mod = static_cast<std::uint16_t>(b + 1);
|
||||
for (unsigned i = 0; i < b; ++i)
|
||||
bits.push_back(c.input(mod));
|
||||
c.out(c.threshold(bits, static_cast<std::uint16_t>(b / 2)));
|
||||
std::vector<std::uint16_t> in(b, 0);
|
||||
for (unsigned i = 0; i < b; i += 2)
|
||||
in[i] = 1;
|
||||
auto got = dpf::arith_garble::eval_pair(c, in.data());
|
||||
return got.opened.empty() ? 0 : got.opened[0];
|
||||
}
|
||||
|
||||
inline std::uint64_t run_chain()
|
||||
{
|
||||
arith_garble::circuit c;
|
||||
auto x = c.input(7);
|
||||
auto y = c.input(7);
|
||||
auto z = c.mul(x, y);
|
||||
for (int i = 0; i < 3; ++i)
|
||||
z = c.mul(z, x);
|
||||
c.out(z);
|
||||
const std::uint16_t in[2] = {2, 3};
|
||||
auto got = dpf::arith_garble::eval_pair(c, in);
|
||||
return got.opened.empty() ? 0 : got.opened[0];
|
||||
}
|
||||
|
||||
inline std::size_t yao_if_bytes(unsigned heavy, unsigned light)
|
||||
{
|
||||
std::vector<std::uint8_t> h(heavy, 1), l(light, 1);
|
||||
auto got = dpf::yao::eval_if(and_n(heavy), and_n(light), 0, 0, h.data(),
|
||||
h.data(), l.data(), l.data());
|
||||
return std::max<std::size_t>(
|
||||
16, (got.stack_blocks + got.extra_blocks) * 16u);
|
||||
}
|
||||
|
||||
inline std::uint64_t run_if(unsigned heavy, unsigned light)
|
||||
{
|
||||
std::vector<std::uint8_t> h0(heavy, 1), h1(heavy, 0), l0(light, 1), l1(light, 1);
|
||||
auto got = dpf::yao::eval_if(and_n(heavy), and_n(light), 1, 0, h0.data(),
|
||||
h1.data(), l0.data(), l1.data());
|
||||
return got.share0.empty()
|
||||
? 0
|
||||
: static_cast<std::uint64_t>(got.share0[0] ^ got.share1[0]);
|
||||
}
|
||||
|
||||
inline std::size_t yao_hot_bytes(unsigned k)
|
||||
{
|
||||
std::vector<dpf::yao::netlist> br;
|
||||
std::vector<std::vector<std::uint8_t>> p0, p1;
|
||||
for (unsigned i = 0; i < k; ++i)
|
||||
{
|
||||
br.push_back(i % 2 == 0 ? and_n(2) : xor2());
|
||||
p0.push_back({1, 1});
|
||||
p1.push_back({0, 0});
|
||||
}
|
||||
auto got = dpf::yao::eval_one_hot(br, 1, 0, p0, p1);
|
||||
return std::max<std::size_t>(
|
||||
16, (got.stack_blocks + got.extra_blocks) * 16u);
|
||||
}
|
||||
|
||||
inline std::uint64_t run_hot(unsigned k)
|
||||
{
|
||||
std::vector<dpf::yao::netlist> br;
|
||||
std::vector<std::vector<std::uint8_t>> p0, p1;
|
||||
for (unsigned i = 0; i < k; ++i)
|
||||
{
|
||||
br.push_back(i % 2 == 0 ? and_n(2) : xor2());
|
||||
p0.push_back({1, static_cast<std::uint8_t>(i & 1u)});
|
||||
p1.push_back({0, 1});
|
||||
}
|
||||
auto got = dpf::yao::eval_one_hot(br, 1, 2, p0, p1);
|
||||
return got.share0.empty()
|
||||
? 0
|
||||
: static_cast<std::uint64_t>(got.share0[0] ^ got.share1[0]);
|
||||
}
|
||||
|
||||
inline std::uint64_t run_flute(unsigned delta, unsigned n_out)
|
||||
{
|
||||
const unsigned rows = 1u << delta;
|
||||
std::vector<std::uint8_t> columns(static_cast<std::size_t>(n_out) * rows);
|
||||
for (unsigned w = 0; w < n_out; ++w)
|
||||
for (unsigned j = 0; j < rows; ++j)
|
||||
columns[static_cast<std::size_t>(w) * rows + j] =
|
||||
static_cast<std::uint8_t>(((j >> (w % delta)) ^ w) & 1u);
|
||||
std::vector<std::uint8_t> bits(delta, 0);
|
||||
bits[0] = 1;
|
||||
if (delta > 2)
|
||||
bits[2] = 1;
|
||||
auto got = dpf::flute::eval_pair(delta, n_out, columns.data(), bits.data());
|
||||
std::uint64_t mix = 0;
|
||||
for (auto b : got.opened)
|
||||
mix = (mix << 1) | b;
|
||||
return mix;
|
||||
}
|
||||
|
||||
inline rss::seed_bundle fixed_bundle()
|
||||
{
|
||||
rss::seed_bundle b{};
|
||||
auto fill = [](rss::seed_block & s, std::uint8_t tag) {
|
||||
auto * p = reinterpret_cast<std::uint8_t *>(&s);
|
||||
for (std::size_t i = 0; i < sizeof(s); ++i)
|
||||
p[i] = static_cast<std::uint8_t>(tag + i * 17u);
|
||||
};
|
||||
fill(b.k01, 1);
|
||||
fill(b.k12, 2);
|
||||
fill(b.k20, 3);
|
||||
return b;
|
||||
}
|
||||
|
||||
inline void shuffle_outbound(unsigned me, unsigned pass, std::size_t n,
|
||||
std::uint8_t * out, std::size_t out_n)
|
||||
{
|
||||
if (out_n != n * sizeof(std::uint64_t))
|
||||
throw std::logic_error("bench shuffle slot");
|
||||
const unsigned order[3] = {2u, 0u, 1u};
|
||||
const unsigned left = order[pass];
|
||||
if (me == left)
|
||||
{
|
||||
std::memset(out, 0, out_n);
|
||||
return;
|
||||
}
|
||||
const auto bundle = fixed_bundle();
|
||||
std::vector<std::uint64_t> column(n);
|
||||
std::iota(column.begin(), column.end(), 0);
|
||||
shuffle::shuffle_party_view<std::uint64_t> held[3];
|
||||
for (unsigned p = 0; p < 3; ++p)
|
||||
{
|
||||
held[p].own.assign(n, 0);
|
||||
held[p].next.assign(n, 0);
|
||||
}
|
||||
std::uint64_t rng = 0xA5A5A5A5A5A5A5A5ull;
|
||||
auto draw = [&] {
|
||||
rng = rng * 6364136223846793005ull + 1u;
|
||||
return rng;
|
||||
};
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
const auto a = draw();
|
||||
const auto b = draw();
|
||||
const auto c = column[i] - a - b;
|
||||
held[0].own[i] = a;
|
||||
held[0].next[i] = b;
|
||||
held[1].own[i] = b;
|
||||
held[1].next[i] = c;
|
||||
held[2].own[i] = c;
|
||||
held[2].next[i] = a;
|
||||
}
|
||||
for (unsigned step = 0; step <= pass; ++step)
|
||||
{
|
||||
const unsigned L = order[step];
|
||||
const unsigned u = shuffle::hidden_u_party(L);
|
||||
const unsigned side = shuffle::hidden_side_party(L);
|
||||
auto u_step = shuffle::shuffle_hidden_pass<std::uint64_t>(
|
||||
u, rss::party_seeds::from_bundle(bundle, u), held[u], 0, L, nullptr);
|
||||
auto s_step = shuffle::shuffle_hidden_pass<std::uint64_t>(side,
|
||||
rss::party_seeds::from_bundle(bundle, side), held[side], 0, L,
|
||||
&u_step.out.data);
|
||||
std::vector<std::uint64_t> side_msg = s_step.out.data;
|
||||
auto l_step = shuffle::shuffle_hidden_pass<std::uint64_t>(L,
|
||||
rss::party_seeds::from_bundle(bundle, L), held[L], 0, L, &side_msg);
|
||||
if (step == pass)
|
||||
{
|
||||
const auto & msg = (me == u) ? u_step.out.data : s_step.out.data;
|
||||
std::memcpy(out, msg.data(), out_n);
|
||||
}
|
||||
held[u] = std::move(u_step.view);
|
||||
held[side] = std::move(s_step.view);
|
||||
held[L] = std::move(l_step.view);
|
||||
}
|
||||
}
|
||||
|
||||
inline std::uint64_t heavy(kind k)
|
||||
{
|
||||
switch (k)
|
||||
{
|
||||
case kind::proj5:
|
||||
return run_proj(5);
|
||||
case kind::proj17:
|
||||
return run_proj(17);
|
||||
case kind::proj64:
|
||||
return run_proj(64);
|
||||
case kind::mul5:
|
||||
return run_mul(5);
|
||||
case kind::mul7:
|
||||
return run_mul(7);
|
||||
case kind::mul11:
|
||||
return run_mul(11);
|
||||
case kind::thresh8:
|
||||
return run_thresh(8);
|
||||
case kind::thresh16:
|
||||
return run_thresh(16);
|
||||
case kind::chain7:
|
||||
return run_chain();
|
||||
case kind::yao_if4:
|
||||
return run_if(4, 2);
|
||||
case kind::yao_if16:
|
||||
return run_if(16, 8);
|
||||
case kind::yao_hot4:
|
||||
return run_hot(4);
|
||||
case kind::yao_hot8:
|
||||
return run_hot(8);
|
||||
case kind::flute2:
|
||||
return run_flute(2, 1);
|
||||
case kind::flute4:
|
||||
return run_flute(4, 1);
|
||||
case kind::flute8:
|
||||
return run_flute(8, 1);
|
||||
case kind::flute4x8:
|
||||
return run_flute(4, 8);
|
||||
default:
|
||||
return 0;
|
||||
}
|
||||
}
|
||||
|
||||
inline std::size_t payload_of(kind k)
|
||||
{
|
||||
switch (k)
|
||||
{
|
||||
case kind::proj5:
|
||||
return proj_bytes(5);
|
||||
case kind::proj17:
|
||||
return proj_bytes(17);
|
||||
case kind::proj64:
|
||||
return proj_bytes(64);
|
||||
case kind::mul5:
|
||||
return mul_bytes(5);
|
||||
case kind::mul7:
|
||||
return mul_bytes(7);
|
||||
case kind::mul11:
|
||||
return mul_bytes(11);
|
||||
case kind::thresh8:
|
||||
return thresh_bytes(8);
|
||||
case kind::thresh16:
|
||||
return thresh_bytes(16);
|
||||
case kind::chain7:
|
||||
return chain_bytes();
|
||||
case kind::yao_if4:
|
||||
return yao_if_bytes(4, 2);
|
||||
case kind::yao_if16:
|
||||
return yao_if_bytes(16, 8);
|
||||
case kind::yao_hot4:
|
||||
return yao_hot_bytes(4);
|
||||
case kind::yao_hot8:
|
||||
return yao_hot_bytes(8);
|
||||
case kind::flute2:
|
||||
case kind::flute4:
|
||||
case kind::flute8:
|
||||
return 8;
|
||||
case kind::flute4x8:
|
||||
return 16;
|
||||
case kind::shuf16:
|
||||
return 16 * sizeof(std::uint64_t);
|
||||
case kind::shuf64:
|
||||
return 64 * sizeof(std::uint64_t);
|
||||
case kind::shuf256:
|
||||
return 256 * sizeof(std::uint64_t);
|
||||
}
|
||||
throw std::invalid_argument("bench: cell");
|
||||
}
|
||||
|
||||
inline protocol::plan peer_plan(std::size_t, kind k)
|
||||
{
|
||||
protocol::composer c(0);
|
||||
auto done = c.bench_peer(static_cast<std::uint32_t>(k), payload_of(k));
|
||||
(void)done;
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
inline protocol::plan ring_plan(std::size_t, kind k)
|
||||
{
|
||||
protocol::composer c(0);
|
||||
auto done = c.bench_ring(static_cast<std::uint32_t>(k), payload_of(k));
|
||||
(void)done;
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
struct cell
|
||||
{
|
||||
const char * name;
|
||||
int parties;
|
||||
kind id;
|
||||
};
|
||||
|
||||
/// @brief The battery. Secret indexes stay on the DPF plans beside these.
|
||||
inline std::vector<cell> battery()
|
||||
{
|
||||
return {
|
||||
{"arith_proj_m5", 2, kind::proj5},
|
||||
{"arith_proj_m17", 2, kind::proj17},
|
||||
{"arith_proj_m64", 2, kind::proj64},
|
||||
{"arith_mul_p5", 2, kind::mul5},
|
||||
{"arith_mul_p7", 2, kind::mul7},
|
||||
{"arith_mul_p11", 2, kind::mul11},
|
||||
{"arith_thresh_b8", 2, kind::thresh8},
|
||||
{"arith_thresh_b16", 2, kind::thresh16},
|
||||
{"arith_chain_mul4", 2, kind::chain7},
|
||||
{"yao_if_4_2", 2, kind::yao_if4},
|
||||
{"yao_if_16_8", 2, kind::yao_if16},
|
||||
{"yao_onehot_k4", 2, kind::yao_hot4},
|
||||
{"yao_onehot_k8", 2, kind::yao_hot8},
|
||||
{"flute_d2", 2, kind::flute2},
|
||||
{"flute_d4", 2, kind::flute4},
|
||||
{"flute_d8", 2, kind::flute8},
|
||||
{"flute_d4_o8", 2, kind::flute4x8},
|
||||
{"shuffle_n16", 3, kind::shuf16},
|
||||
{"shuffle_n64", 3, kind::shuf64},
|
||||
{"shuffle_n256", 3, kind::shuf256},
|
||||
};
|
||||
}
|
||||
|
||||
inline protocol::plan plan_for(std::size_t party, kind id)
|
||||
{
|
||||
if (is_shuffle(id))
|
||||
return detail::ring_plan(party, id);
|
||||
return detail::peer_plan(party, id);
|
||||
}
|
||||
|
||||
/// @brief Local half of a cell. Party 0 emits the payload. Party 1 applies it.
|
||||
/// Every shuffle party emits its own ring slot.
|
||||
inline void run_cell(std::uint32_t opcode, std::size_t party, std::uint32_t aux,
|
||||
std::uint8_t * out, std::size_t out_n)
|
||||
{
|
||||
if (out == nullptr && out_n != 0)
|
||||
throw std::invalid_argument("bench cell buffer");
|
||||
const auto k = kind_of(aux);
|
||||
if (is_shuffle(k))
|
||||
{
|
||||
if (opcode != protocol::opcodes::bench_emit)
|
||||
return;
|
||||
detail::shuffle_outbound(static_cast<unsigned>(party), pass_of(aux),
|
||||
shuffle_n(k), out, out_n);
|
||||
return;
|
||||
}
|
||||
if (opcode == protocol::opcodes::bench_apply)
|
||||
{
|
||||
if (party != 1)
|
||||
{
|
||||
if (out_n != 0)
|
||||
out[0] = 0;
|
||||
return;
|
||||
}
|
||||
const auto mix = detail::heavy(k);
|
||||
if (out_n >= sizeof(mix))
|
||||
std::memcpy(out, &mix, sizeof(mix));
|
||||
return;
|
||||
}
|
||||
if (party != 0)
|
||||
{
|
||||
if (out_n != 0)
|
||||
std::memset(out, 0, out_n);
|
||||
return;
|
||||
}
|
||||
detail::paint(out, out_n, detail::heavy(k));
|
||||
}
|
||||
|
||||
} // namespace bench
|
||||
} // namespace dpf
|
||||
|
||||
#endif
|
||||
|
|
@ -15,7 +15,7 @@
|
|||
/// @author Ryan Henry <ryan.henry@ucalgary.ca>
|
||||
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
|
||||
/// @license Released under a GNU General Public v2.0 (GPLv2) license;
|
||||
/// see [LICENSE.md](@ref license) for details.
|
||||
/// see LICENSE.md for details.
|
||||
|
||||
#ifndef LIBDPF_INCLUDE_DPF_BIT_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_BIT_HPP__
|
||||
|
|
@ -38,6 +38,14 @@ namespace dpf
|
|||
{
|
||||
|
||||
/// @brief binary type whose representation can be packed into one bit
|
||||
/// @note Not a DPF domain. `msb_of<dpf::bit>` does not compile. A 1-bit
|
||||
/// index is `dpf::modint<1>` or `dpf::xint<1>`. Leaf `+` and `-` are
|
||||
/// XOR. Lanes pack low-bit first.
|
||||
/// @see dpf::twobit
|
||||
/// @see dpf::nyble
|
||||
/// @see dpf::modint
|
||||
/// @see dpf::xint
|
||||
/// @see output_types
|
||||
enum bit : bool
|
||||
{
|
||||
zero = false, ///< `0`, `false`, "unset", "off"
|
||||
|
|
@ -225,6 +233,10 @@ struct bitlength_of_output<dpf::bit, NodeT>
|
|||
template <>
|
||||
struct is_packed_subbyte<dpf::bit> : std::true_type {};
|
||||
|
||||
/// @brief Packed bits add by XOR (`bit::one + bit::one == bit::zero`).
|
||||
template <>
|
||||
struct has_characteristic_two<dpf::bit> : std::true_type {};
|
||||
|
||||
template <>
|
||||
struct packed_lane_bits<dpf::bit>
|
||||
: public std::integral_constant<std::size_t, 1> {};
|
||||
|
|
|
|||
|
|
@ -43,6 +43,7 @@ template <typename Word>
|
|||
HEDLEY_NO_THROW
|
||||
constexpr void check_one_bit(Word mask) noexcept
|
||||
{
|
||||
(void)mask;
|
||||
#if defined(__GNUC__) || defined(__clang__)
|
||||
if (__builtin_is_constant_evaluated()) return;
|
||||
#endif
|
||||
|
|
|
|||
77
include/dpf/bit_inject.hpp
Normal file
77
include/dpf/bit_inject.hpp
Normal file
|
|
@ -0,0 +1,77 @@
|
|||
/// @file dpf/bit_inject.hpp
|
||||
/// @brief Boolean bit × arithmetic value (2PC bit_mul and 3PC RSS injection).
|
||||
#ifndef LIBDPF_INCLUDE_DPF_BIT_INJECT_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_BIT_INJECT_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <utility>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/beaver.hpp"
|
||||
#include "dpf/rss_seed.hpp"
|
||||
#include "dpf/secret_share.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace bit_inject
|
||||
{
|
||||
|
||||
/// @brief 2PC: `b * x` via a beaver bit_mul on an existing session.
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
beavers::wire<Ring> inject2(beavers::session<Ring> & s,
|
||||
beavers::wire<Ring> bit_wire, beavers::wire<Ring> arith_wire)
|
||||
{
|
||||
return s.bit_mul(bit_wire, arith_wire);
|
||||
}
|
||||
|
||||
/// @brief Local y-factor of RSS bit injection before ring send.
|
||||
/// @details Party holds RSS bit `(b_own, b_next)` and arithmetic `(x_own, x_next)`.
|
||||
/// Local contribution mirrors RSS mul with the bit as a 0/1 factor.
|
||||
template <typename T>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
T rss_inject_local(const rss::party_seeds & seeds, T b_own, T b_next,
|
||||
T x_own, T x_next, std::uint64_t index)
|
||||
{
|
||||
// Treat bit components as ring elements in {0,1}.
|
||||
return rss::rss_mul_local(seeds, b_own, b_next, x_own, x_next, index);
|
||||
}
|
||||
|
||||
/// @brief After ring refresh: party receives `y_prev` and forms RSS `(y_own, y_prev)`
|
||||
/// wait — standard RSS refresh: send y_own to next, receive from prev,
|
||||
/// store `(y_own, y_from_prev)`? Actually ABY3: party i holds y_i after
|
||||
/// local mul; sends y_i to party i-1; ends with (y_i, y_{i+1}).
|
||||
template <typename T>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::pair<T, T> rss_refresh(T y_own, T y_from_next)
|
||||
{
|
||||
return {y_own, y_from_next};
|
||||
}
|
||||
|
||||
/// @brief Cleartext identity: `(b0⊕b1) * (x0+x1)`.
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
Ring inject_clear(std::uint8_t b0, std::uint8_t b1, Ring x0, Ring x1)
|
||||
{
|
||||
const Ring b = static_cast<Ring>((b0 ^ b1) & 1u);
|
||||
return static_cast<Ring>(b * (x0 + x1));
|
||||
}
|
||||
|
||||
/// @brief Boolean RSS AND local factor (GF(2)).
|
||||
inline std::uint8_t rss_and_local(const rss::party_seeds & seeds,
|
||||
std::uint8_t a_own, std::uint8_t a_next, std::uint8_t b_own,
|
||||
std::uint8_t b_next, std::uint64_t index)
|
||||
{
|
||||
const std::uint8_t cross = static_cast<std::uint8_t>(
|
||||
(a_own & b_own) ^ (a_own & b_next) ^ (a_next & b_own));
|
||||
const std::uint8_t mask = static_cast<std::uint8_t>(
|
||||
rss::zero_share<std::uint8_t>(seeds, index) & 1u);
|
||||
return static_cast<std::uint8_t>(cross ^ mask);
|
||||
}
|
||||
|
||||
} // namespace bit_inject
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_BIT_INJECT_HPP__
|
||||
781
include/dpf/bitmore_mod.hpp
Normal file
781
include/dpf/bitmore_mod.hpp
Normal file
|
|
@ -0,0 +1,781 @@
|
|||
/// @file dpf/bitmore_mod.hpp
|
||||
/// @brief Byte-slot reduction for BitMore moduli through 255.
|
||||
/// @details Hafiz and Henry (PoPETs 2019 §5.3) fold each server's DPF bits
|
||||
/// into an integer and reduce modulo the server count. When that
|
||||
/// count is not a power of two the integer does not fit in a byte
|
||||
/// for long, so each byte is only partially reduced until the end.
|
||||
///
|
||||
/// A partial step reads the high nibble `h`. Every byte with that
|
||||
/// nibble is at least `16*h`, so `floor(16*h / M) * M` is the largest
|
||||
/// multiple of `M` that is safe to subtract. `pshufb` selects it and
|
||||
/// `sub_epi8` removes it. The byte stays congruent modulo `M` and is
|
||||
/// at most `partial_bound`. For moduli 121..127 that ceiling is 128
|
||||
/// or more, so a second partial step is what leaves `stable_bound`
|
||||
/// (at most 127). `resume_bound` is the ceiling the accumulator
|
||||
/// actually resumes from: one step through modulus 120 and 128, two
|
||||
/// steps for 121..127.
|
||||
///
|
||||
/// `add_budget = 255 - resume_bound` is how much can still be added
|
||||
/// before a byte might reach 256. `shift_budget` is how many
|
||||
/// `2*acc + bit` insertions fit in that slack. Both stay positive
|
||||
/// through modulus 128, which is as far as a byte can hold two
|
||||
/// resumed slots or one more bit. `partial_reduce` and `full_reduce`
|
||||
/// themselves stay correct through 255: the two nibble residues sum
|
||||
/// to at most 255 and to less than `2*M`, and one unsigned compare
|
||||
/// subtracts `M`. Above 128 there is no slack left to defer that
|
||||
/// correction across another add.
|
||||
///
|
||||
/// `bitmore_mod<M, 16>` is the same idea on 16-bit lanes, still on
|
||||
/// AVX2. The top nibble (bits 12..15) selects `floor(4096*h / M)*M`
|
||||
/// through two `pshufb`s, one per byte of that multiple, and
|
||||
/// `sub_epi16` removes it. The accumulator has slack through modulus
|
||||
/// 32768 (`resume_bound` at most 32767, so one more bit fits in the
|
||||
/// lane). A full reduction sums the four nibble residues, which fit
|
||||
/// in the lane for every modulus through 65535, then subtracts `M`
|
||||
/// up to three times.
|
||||
/// @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_DPF_BITMORE_MOD_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_BITMORE_MOD_HPP__
|
||||
|
||||
#include <array>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <type_traits>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
#include "simde/simde/x86/avx2.h"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace bitmore_detail
|
||||
{
|
||||
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg zero() noexcept
|
||||
{
|
||||
if constexpr (std::is_same_v<Reg, simde__m256i>)
|
||||
return simde_mm256_setzero_si256();
|
||||
else
|
||||
return simde_mm_setzero_si128();
|
||||
}
|
||||
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg load_lut(const std::array<unsigned char, 16> & lut) noexcept
|
||||
{
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
{
|
||||
simde__m128i table;
|
||||
std::memcpy(&table, lut.data(), 16);
|
||||
return table;
|
||||
}
|
||||
else
|
||||
{
|
||||
alignas(32) unsigned char both[32];
|
||||
std::memcpy(both, lut.data(), 16);
|
||||
std::memcpy(both + 16, lut.data(), 16);
|
||||
simde__m256i table;
|
||||
std::memcpy(&table, both, 32);
|
||||
return table;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg high_nibble(Reg x) noexcept
|
||||
{
|
||||
// `srli_epi16` moves the neighbouring byte's low nibble into bits 4..7.
|
||||
// Masking with `0x0f` leaves this byte's own high nibble.
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
{
|
||||
const auto m = simde_mm_set1_epi8(0x0f);
|
||||
return simde_mm_and_si128(simde_mm_srli_epi16(x, 4), m);
|
||||
}
|
||||
else
|
||||
{
|
||||
const auto m = simde_mm256_set1_epi8(0x0f);
|
||||
return simde_mm256_and_si256(simde_mm256_srli_epi16(x, 4), m);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg low_nibble(Reg x) noexcept
|
||||
{
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
return simde_mm_and_si128(x, simde_mm_set1_epi8(0x0f));
|
||||
else
|
||||
return simde_mm256_and_si256(x, simde_mm256_set1_epi8(0x0f));
|
||||
}
|
||||
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg shuffle(Reg table, Reg index) noexcept
|
||||
{
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
return simde_mm_shuffle_epi8(table, index);
|
||||
else
|
||||
return simde_mm256_shuffle_epi8(table, index);
|
||||
}
|
||||
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg add_bytes(Reg a, Reg b) noexcept
|
||||
{
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
return simde_mm_add_epi8(a, b);
|
||||
else
|
||||
return simde_mm256_add_epi8(a, b);
|
||||
}
|
||||
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg sub_bytes(Reg a, Reg b) noexcept
|
||||
{
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
return simde_mm_sub_epi8(a, b);
|
||||
else
|
||||
return simde_mm256_sub_epi8(a, b);
|
||||
}
|
||||
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg and_bytes(Reg a, Reg b) noexcept
|
||||
{
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
return simde_mm_and_si128(a, b);
|
||||
else
|
||||
return simde_mm256_and_si256(a, b);
|
||||
}
|
||||
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg xor_bytes(Reg a, Reg b) noexcept
|
||||
{
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
return simde_mm_xor_si128(a, b);
|
||||
else
|
||||
return simde_mm256_xor_si256(a, b);
|
||||
}
|
||||
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg splat_epi8(unsigned char v) noexcept
|
||||
{
|
||||
const auto s = static_cast<int8_t>(v);
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
return simde_mm_set1_epi8(s);
|
||||
else
|
||||
return simde_mm256_set1_epi8(s);
|
||||
}
|
||||
|
||||
/// @brief Unsigned `a > b` per byte. `cmpgt_epi8` is signed; XOR `0x80` maps
|
||||
/// unsigned order onto that signed order.
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg cmpgt_epu8(Reg a, Reg b) noexcept
|
||||
{
|
||||
const auto bias = splat_epi8<Reg>(0x80);
|
||||
const auto aa = xor_bytes<Reg>(a, bias);
|
||||
const auto bb = xor_bytes<Reg>(b, bias);
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
return simde_mm_cmpgt_epi8(aa, bb);
|
||||
else
|
||||
return simde_mm256_cmpgt_epi8(aa, bb);
|
||||
}
|
||||
|
||||
/// @brief Shift each byte left by 1 and set bit 0 from `bit`.
|
||||
/// @details `slli_epi16` spills bit 7 into the next byte of the 16-bit lane.
|
||||
/// Clearing bit 0 afterwards drops that spill. Bit 7 of the odd byte
|
||||
/// shifts out of the lane, which is the byte-local shift.
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg shift_in_bit(Reg acc, Reg bit) noexcept
|
||||
{
|
||||
const auto one = splat_epi8<Reg>(1);
|
||||
const auto keep = splat_epi8<Reg>(0xfe);
|
||||
bit = and_bytes<Reg>(bit, one);
|
||||
Reg shifted;
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
shifted = simde_mm_and_si128(simde_mm_slli_epi16(acc, 1), keep);
|
||||
else
|
||||
shifted = simde_mm256_and_si256(simde_mm256_slli_epi16(acc, 1), keep);
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
return simde_mm_or_si128(shifted, bit);
|
||||
else
|
||||
return simde_mm256_or_si256(shifted, bit);
|
||||
}
|
||||
|
||||
template <unsigned Modulus, unsigned LaneBits = 8>
|
||||
HEDLEY_CONST
|
||||
constexpr unsigned partial_bound() noexcept
|
||||
{
|
||||
static_assert(LaneBits == 8u || LaneBits == 16u, "bitmore lane width is 8 or 16");
|
||||
constexpr unsigned block = LaneBits == 8u ? 16u : 4096u;
|
||||
constexpr unsigned tail = block - 1u;
|
||||
unsigned u = 0;
|
||||
for (unsigned h = 0; h < 16u; ++h)
|
||||
{
|
||||
const unsigned q = (block * h / Modulus) * Modulus;
|
||||
const unsigned top = block * h + tail - q;
|
||||
if (top > u)
|
||||
u = top;
|
||||
}
|
||||
return u;
|
||||
}
|
||||
|
||||
template <unsigned Modulus, unsigned LaneBits = 8>
|
||||
HEDLEY_CONST
|
||||
constexpr unsigned stable_bound() noexcept
|
||||
{
|
||||
constexpr unsigned block = LaneBits == 8u ? 16u : 4096u;
|
||||
const unsigned lim = partial_bound<Modulus, LaneBits>();
|
||||
unsigned u = 0;
|
||||
for (unsigned h = 0; h < 16u; ++h)
|
||||
{
|
||||
const unsigned lo = block * h;
|
||||
if (lo > lim)
|
||||
break;
|
||||
unsigned hi = lo + block - 1u;
|
||||
if (hi > lim)
|
||||
hi = lim;
|
||||
const unsigned q = (lo / Modulus) * Modulus;
|
||||
const unsigned top = hi - q;
|
||||
if (top > u)
|
||||
u = top;
|
||||
}
|
||||
return u;
|
||||
}
|
||||
|
||||
template <unsigned Modulus>
|
||||
HEDLEY_CONST
|
||||
constexpr unsigned shift_budget(unsigned bound, unsigned capacity = 256u) noexcept
|
||||
{
|
||||
unsigned k = 0;
|
||||
unsigned span = bound + 1u;
|
||||
const unsigned half = capacity >> 1;
|
||||
while (span <= half)
|
||||
{
|
||||
span *= 2u;
|
||||
++k;
|
||||
}
|
||||
return k;
|
||||
}
|
||||
|
||||
template <unsigned Modulus>
|
||||
HEDLEY_CONST
|
||||
constexpr std::array<unsigned char, 16> partial_lut() noexcept
|
||||
{
|
||||
std::array<unsigned char, 16> lut{};
|
||||
for (unsigned h = 0; h < 16u; ++h)
|
||||
lut[h] = static_cast<unsigned char>((16u * h / Modulus) * Modulus);
|
||||
return lut;
|
||||
}
|
||||
|
||||
template <unsigned Modulus>
|
||||
HEDLEY_CONST
|
||||
constexpr std::array<unsigned char, 16> low_residue_lut() noexcept
|
||||
{
|
||||
std::array<unsigned char, 16> lut{};
|
||||
for (unsigned n = 0; n < 16u; ++n)
|
||||
lut[n] = static_cast<unsigned char>(n % Modulus);
|
||||
return lut;
|
||||
}
|
||||
|
||||
template <unsigned Modulus>
|
||||
HEDLEY_CONST
|
||||
constexpr std::array<unsigned char, 16> high_residue_lut() noexcept
|
||||
{
|
||||
std::array<unsigned char, 16> lut{};
|
||||
for (unsigned h = 0; h < 16u; ++h)
|
||||
lut[h] = static_cast<unsigned char>((16u * h) % Modulus);
|
||||
return lut;
|
||||
}
|
||||
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg splat_epi16(unsigned v) noexcept
|
||||
{
|
||||
const auto s = static_cast<std::int16_t>(static_cast<std::uint16_t>(v));
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
return simde_mm_set1_epi16(s);
|
||||
else
|
||||
return simde_mm256_set1_epi16(s);
|
||||
}
|
||||
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg add_epi16(Reg a, Reg b) noexcept
|
||||
{
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
return simde_mm_add_epi16(a, b);
|
||||
else
|
||||
return simde_mm256_add_epi16(a, b);
|
||||
}
|
||||
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg sub_epi16(Reg a, Reg b) noexcept
|
||||
{
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
return simde_mm_sub_epi16(a, b);
|
||||
else
|
||||
return simde_mm256_sub_epi16(a, b);
|
||||
}
|
||||
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg srli_epi16(Reg a, int imm) noexcept
|
||||
{
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
return simde_mm_srli_epi16(a, imm);
|
||||
else
|
||||
return simde_mm256_srli_epi16(a, imm);
|
||||
}
|
||||
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg slli_epi16(Reg a, int imm) noexcept
|
||||
{
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
return simde_mm_slli_epi16(a, imm);
|
||||
else
|
||||
return simde_mm256_slli_epi16(a, imm);
|
||||
}
|
||||
|
||||
/// @brief Unsigned `a > b` per 16-bit lane.
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg cmpgt_epu16(Reg a, Reg b) noexcept
|
||||
{
|
||||
const auto bias = splat_epi16<Reg>(0x8000u);
|
||||
const auto aa = xor_bytes<Reg>(a, bias);
|
||||
const auto bb = xor_bytes<Reg>(b, bias);
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
return simde_mm_cmpgt_epi16(aa, bb);
|
||||
else
|
||||
return simde_mm256_cmpgt_epi16(aa, bb);
|
||||
}
|
||||
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg or_bytes(Reg a, Reg b) noexcept
|
||||
{
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
return simde_mm_or_si128(a, b);
|
||||
else
|
||||
return simde_mm256_or_si256(a, b);
|
||||
}
|
||||
|
||||
/// @brief Look up a 16-bit table entry. `idx` holds the nibble in the low byte
|
||||
/// of each lane and zero in the high byte, so `pshufb` writes `lut[h]`
|
||||
/// into the low byte and `lut[0]` into the high byte. `lut[0]` is 0.
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg lookup_u16(Reg idx, const std::array<unsigned char, 16> & lo,
|
||||
const std::array<unsigned char, 16> & hi) noexcept
|
||||
{
|
||||
const auto lo_v = shuffle<Reg>(load_lut<Reg>(lo), idx);
|
||||
const auto hi_v = slli_epi16<Reg>(shuffle<Reg>(load_lut<Reg>(hi), idx), 8);
|
||||
return or_bytes<Reg>(lo_v, hi_v);
|
||||
}
|
||||
|
||||
constexpr std::array<unsigned char, 16> u16_lo(const std::array<std::uint16_t, 16> & v) noexcept
|
||||
{
|
||||
std::array<unsigned char, 16> out{};
|
||||
for (unsigned i = 0; i < 16u; ++i)
|
||||
out[i] = static_cast<unsigned char>(v[i] & 0xffu);
|
||||
return out;
|
||||
}
|
||||
|
||||
constexpr std::array<unsigned char, 16> u16_hi(const std::array<std::uint16_t, 16> & v) noexcept
|
||||
{
|
||||
std::array<unsigned char, 16> out{};
|
||||
for (unsigned i = 0; i < 16u; ++i)
|
||||
out[i] = static_cast<unsigned char>(v[i] >> 8);
|
||||
return out;
|
||||
}
|
||||
|
||||
template <unsigned Modulus>
|
||||
HEDLEY_CONST
|
||||
constexpr std::array<std::uint16_t, 16> wide_partial_lut() noexcept
|
||||
{
|
||||
std::array<std::uint16_t, 16> lut{};
|
||||
for (unsigned h = 0; h < 16u; ++h)
|
||||
lut[h] = static_cast<std::uint16_t>((4096u * h / Modulus) * Modulus);
|
||||
return lut;
|
||||
}
|
||||
|
||||
template <unsigned Modulus, unsigned Place>
|
||||
HEDLEY_CONST
|
||||
constexpr std::array<std::uint16_t, 16> wide_residue_lut() noexcept
|
||||
{
|
||||
std::array<std::uint16_t, 16> lut{};
|
||||
for (unsigned n = 0; n < 16u; ++n)
|
||||
lut[n] = static_cast<std::uint16_t>(
|
||||
(static_cast<std::uint32_t>(Place) * n) % Modulus);
|
||||
return lut;
|
||||
}
|
||||
|
||||
/// @brief `acc = 2*acc + bit0` inside each 16-bit lane. `slli_epi16` does not
|
||||
/// cross lanes.
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
Reg shift_in_bit16(Reg acc, Reg bit) noexcept
|
||||
{
|
||||
const auto one = splat_epi16<Reg>(1u);
|
||||
bit = and_bytes<Reg>(bit, one);
|
||||
if constexpr (std::is_same_v<Reg, simde__m128i>)
|
||||
return simde_mm_or_si128(simde_mm_slli_epi16(acc, 1), bit);
|
||||
else
|
||||
return simde_mm256_or_si256(simde_mm256_slli_epi16(acc, 1), bit);
|
||||
}
|
||||
|
||||
} // namespace bitmore_detail
|
||||
|
||||
/// @brief Partial and full reduction of `LaneBits`-wide slots modulo `Modulus`.
|
||||
/// @tparam Modulus server count. `2` through `255` for bytes, `2` through `65535` for 16-bit lanes.
|
||||
/// @tparam LaneBits `8` (one byte per slot) or `16` (one AVX2 `epi16` lane per slot)
|
||||
template <unsigned Modulus, unsigned LaneBits = 8>
|
||||
struct bitmore_mod
|
||||
{
|
||||
static_assert(LaneBits == 8u || LaneBits == 16u,
|
||||
"bitmore lane width is 8 or 16");
|
||||
static_assert(Modulus >= 2u, "bitmore modulus must be at least 2");
|
||||
static_assert(LaneBits == 16u || Modulus <= 255u,
|
||||
"bitmore byte reduction: modulus must be in 2..255");
|
||||
static_assert(LaneBits == 8u || Modulus <= 65535u,
|
||||
"bitmore 16-bit reduction: modulus must be in 2..65535");
|
||||
|
||||
static constexpr unsigned modulus = Modulus;
|
||||
static constexpr unsigned lane_bits = LaneBits;
|
||||
static constexpr unsigned slot_max = LaneBits == 8u ? 255u : 65535u;
|
||||
static constexpr unsigned nibble_block = LaneBits == 8u ? 16u : 4096u;
|
||||
static constexpr unsigned nibble_shift = LaneBits == 8u ? 4u : 12u;
|
||||
|
||||
/// @brief Largest slot one top-nibble partial reduction can leave.
|
||||
static constexpr unsigned partial_bound = bitmore_detail::partial_bound<Modulus, LaneBits>();
|
||||
|
||||
/// @brief Largest slot a second partial step can leave, starting from `partial_bound`.
|
||||
static constexpr unsigned stable_bound = bitmore_detail::stable_bound<Modulus, LaneBits>();
|
||||
|
||||
/// @brief Ceiling the accumulator resumes from. One step when that already
|
||||
/// fits in the lower half of the lane; otherwise the second step.
|
||||
static constexpr unsigned resume_bound =
|
||||
partial_bound <= (slot_max >> 1) ? partial_bound : stable_bound;
|
||||
|
||||
/// @brief How much can be added to a resumed slot before it might reach `slot_max + 1`.
|
||||
static constexpr unsigned add_budget = slot_max - resume_bound;
|
||||
|
||||
/// @brief Bit insertions that fit in `add_budget` after resuming.
|
||||
static constexpr unsigned shift_budget =
|
||||
bitmore_detail::shift_budget<Modulus>(resume_bound, slot_max + 1u);
|
||||
|
||||
/// @brief `x - floor(block*(x>>shift) / M) * M` inside one slot.
|
||||
/// \complexity One division of a nibble index.
|
||||
HEDLEY_CONST
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static constexpr unsigned partial_reduce_slot(unsigned x) noexcept
|
||||
{
|
||||
x &= slot_max;
|
||||
const unsigned h = x >> nibble_shift;
|
||||
const unsigned q = (nibble_block * h / Modulus) * Modulus;
|
||||
return x - q;
|
||||
}
|
||||
|
||||
/// @brief Byte-slot name for `partial_reduce_slot`.
|
||||
HEDLEY_CONST
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static constexpr unsigned partial_reduce_byte(unsigned x) noexcept
|
||||
{
|
||||
static_assert(LaneBits == 8u, "partial_reduce_byte is the 8-bit slot");
|
||||
return partial_reduce_slot(x);
|
||||
}
|
||||
|
||||
/// @brief `x mod Modulus` for one slot.
|
||||
/// \complexity One remainder of a slot.
|
||||
HEDLEY_CONST
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static constexpr unsigned full_reduce_slot(unsigned x) noexcept
|
||||
{
|
||||
return (x & slot_max) % Modulus;
|
||||
}
|
||||
|
||||
/// @brief Byte-slot name for `full_reduce_slot`.
|
||||
HEDLEY_CONST
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static constexpr unsigned full_reduce_byte(unsigned x) noexcept
|
||||
{
|
||||
static_assert(LaneBits == 8u, "full_reduce_byte is the 8-bit slot");
|
||||
return full_reduce_slot(x);
|
||||
}
|
||||
|
||||
/// @brief Partial-reduce every byte of `x`.
|
||||
/// @details The high nibble selects `floor(16*h / M) * M`. Subtracting it
|
||||
/// preserves the residue and leaves a byte of at most `partial_bound`.
|
||||
/// \complexity One `pshufb` and one `sub_epi8` per register.
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_PURE
|
||||
static Reg partial_reduce(Reg x) noexcept
|
||||
{
|
||||
static_assert(std::is_same_v<Reg, simde__m128i> || std::is_same_v<Reg, simde__m256i>,
|
||||
"bitmore partial_reduce: register must be simde__m128i or simde__m256i");
|
||||
if constexpr (LaneBits == 8u)
|
||||
{
|
||||
constexpr auto lut = bitmore_detail::partial_lut<Modulus>();
|
||||
const auto q = bitmore_detail::shuffle<Reg>(
|
||||
bitmore_detail::load_lut<Reg>(lut),
|
||||
bitmore_detail::high_nibble<Reg>(x));
|
||||
return bitmore_detail::sub_bytes<Reg>(x, q);
|
||||
}
|
||||
else
|
||||
{
|
||||
constexpr auto lut = bitmore_detail::wide_partial_lut<Modulus>();
|
||||
constexpr auto lo = bitmore_detail::u16_lo(lut);
|
||||
constexpr auto hi = bitmore_detail::u16_hi(lut);
|
||||
const auto idx = bitmore_detail::srli_epi16<Reg>(x, 12);
|
||||
const auto q = bitmore_detail::lookup_u16<Reg>(idx, lo, hi);
|
||||
return bitmore_detail::sub_epi16<Reg>(x, q);
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Fully reduce every byte of `x` into `0 .. Modulus-1`.
|
||||
/// @details `pshufb` maps the low nibble to itself modulo `M` and the high
|
||||
/// nibble to `(16*h) mod M`. The sum is less than `2*M`, so one
|
||||
/// compare subtracts `M` where the sum is still too big.
|
||||
/// \complexity Two `pshufb`s, one byte add, one compare, one byte subtract.
|
||||
template <typename Reg>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_PURE
|
||||
static Reg full_reduce(Reg x) noexcept
|
||||
{
|
||||
static_assert(std::is_same_v<Reg, simde__m128i> || std::is_same_v<Reg, simde__m256i>,
|
||||
"bitmore full_reduce: register must be simde__m128i or simde__m256i");
|
||||
if constexpr (LaneBits == 8u)
|
||||
{
|
||||
constexpr auto lo_lut = bitmore_detail::low_residue_lut<Modulus>();
|
||||
constexpr auto hi_lut = bitmore_detail::high_residue_lut<Modulus>();
|
||||
const auto lo = bitmore_detail::shuffle<Reg>(
|
||||
bitmore_detail::load_lut<Reg>(lo_lut),
|
||||
bitmore_detail::low_nibble<Reg>(x));
|
||||
const auto hi = bitmore_detail::shuffle<Reg>(
|
||||
bitmore_detail::load_lut<Reg>(hi_lut),
|
||||
bitmore_detail::high_nibble<Reg>(x));
|
||||
const auto sum = bitmore_detail::add_bytes<Reg>(lo, hi);
|
||||
// Sum of the two residues is at most 255 and less than `2*M`, so one
|
||||
// subtraction finishes the byte. The compare is unsigned: for M > 120
|
||||
// the sum can exceed 127.
|
||||
const auto limit = bitmore_detail::splat_epi8<Reg>(
|
||||
static_cast<unsigned char>(Modulus - 1u));
|
||||
const auto ge = bitmore_detail::cmpgt_epu8<Reg>(sum, limit);
|
||||
const auto corr = bitmore_detail::and_bytes<Reg>(
|
||||
ge, bitmore_detail::splat_epi8<Reg>(static_cast<unsigned char>(Modulus)));
|
||||
return bitmore_detail::sub_bytes<Reg>(sum, corr);
|
||||
}
|
||||
else
|
||||
{
|
||||
// Four nibble residues, each `< M`. Their sum fits in a 16-bit lane
|
||||
// for every modulus through 65535 and is less than `4*M`, so three
|
||||
// conditional subtractions finish the lane.
|
||||
constexpr auto r0 = bitmore_detail::wide_residue_lut<Modulus, 1u>();
|
||||
constexpr auto r1 = bitmore_detail::wide_residue_lut<Modulus, 16u>();
|
||||
constexpr auto r2 = bitmore_detail::wide_residue_lut<Modulus, 256u>();
|
||||
constexpr auto r3 = bitmore_detail::wide_residue_lut<Modulus, 4096u>();
|
||||
constexpr auto r0_lo = bitmore_detail::u16_lo(r0);
|
||||
constexpr auto r0_hi = bitmore_detail::u16_hi(r0);
|
||||
constexpr auto r1_lo = bitmore_detail::u16_lo(r1);
|
||||
constexpr auto r1_hi = bitmore_detail::u16_hi(r1);
|
||||
constexpr auto r2_lo = bitmore_detail::u16_lo(r2);
|
||||
constexpr auto r2_hi = bitmore_detail::u16_hi(r2);
|
||||
constexpr auto r3_lo = bitmore_detail::u16_lo(r3);
|
||||
constexpr auto r3_hi = bitmore_detail::u16_hi(r3);
|
||||
const auto nib = bitmore_detail::splat_epi16<Reg>(0x000fu);
|
||||
const auto n0 = bitmore_detail::and_bytes<Reg>(x, nib);
|
||||
const auto n1 = bitmore_detail::and_bytes<Reg>(bitmore_detail::srli_epi16<Reg>(x, 4), nib);
|
||||
const auto n2 = bitmore_detail::and_bytes<Reg>(bitmore_detail::srli_epi16<Reg>(x, 8), nib);
|
||||
const auto n3 = bitmore_detail::srli_epi16<Reg>(x, 12);
|
||||
auto sum = bitmore_detail::lookup_u16<Reg>(n0, r0_lo, r0_hi);
|
||||
sum = bitmore_detail::add_epi16<Reg>(sum, bitmore_detail::lookup_u16<Reg>(n1, r1_lo, r1_hi));
|
||||
sum = bitmore_detail::add_epi16<Reg>(sum, bitmore_detail::lookup_u16<Reg>(n2, r2_lo, r2_hi));
|
||||
sum = bitmore_detail::add_epi16<Reg>(sum, bitmore_detail::lookup_u16<Reg>(n3, r3_lo, r3_hi));
|
||||
const auto limit = bitmore_detail::splat_epi16<Reg>(Modulus - 1u);
|
||||
const auto modv = bitmore_detail::splat_epi16<Reg>(Modulus);
|
||||
for (int step = 0; step < 3; ++step)
|
||||
{
|
||||
const auto ge = bitmore_detail::cmpgt_epu16<Reg>(sum, limit);
|
||||
const auto corr = bitmore_detail::and_bytes<Reg>(ge, modv);
|
||||
sum = bitmore_detail::sub_epi16<Reg>(sum, corr);
|
||||
}
|
||||
return sum;
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
/// @brief Running slot sum modulo `Modulus`, partial-reduced on a budget.
|
||||
/// @details `bound()` is a proven upper bound on every slot. An add whose
|
||||
/// worst-case total would pass the lane maximum partial-reduces the
|
||||
/// accumulator, then the addend. A modulus whose first partial step
|
||||
/// still sets the high bit takes a second step before two slots fit
|
||||
/// again. `insert_bit` is the MSB-first fold `acc = 2*acc + bit`.
|
||||
/// `reduced()` is the full per-slot residue. The byte accumulator
|
||||
/// stops at modulus 128 and the 16-bit accumulator at 32768; past
|
||||
/// that a resumed slot no longer fits next to another or under a
|
||||
/// shift. `partial_reduce` and `full_reduce` cover the wider ranges.
|
||||
/// @tparam Modulus server count, `2` through `128` for bytes and `32768` for 16-bit lanes
|
||||
/// @tparam Reg `simde__m128i` or `simde__m256i`
|
||||
/// @tparam LaneBits `8` or `16`
|
||||
template <unsigned Modulus, typename Reg, unsigned LaneBits = 8>
|
||||
class bitmore_accumulator
|
||||
{
|
||||
using mod = bitmore_mod<Modulus, LaneBits>;
|
||||
|
||||
static_assert(std::is_same_v<Reg, simde__m128i> || std::is_same_v<Reg, simde__m256i>,
|
||||
"bitmore_accumulator: register must be simde__m128i or simde__m256i");
|
||||
static_assert((LaneBits == 8u && Modulus <= 128u) || (LaneBits == 16u && Modulus <= 32768u),
|
||||
"bitmore_accumulator: no slack past modulus 128 in a byte or 32768 in a 16-bit slot");
|
||||
static_assert(mod::stable_bound <= (mod::slot_max >> 1),
|
||||
"bitmore_accumulator: second partial step must leave room for one bit");
|
||||
static_assert(mod::stable_bound * 2u <= mod::slot_max,
|
||||
"bitmore_accumulator: two resumed slots must fit in one lane");
|
||||
static_assert(mod::shift_budget >= 1u,
|
||||
"bitmore_accumulator: at least one bit insertion after resuming");
|
||||
|
||||
public:
|
||||
/// \complexity One zeroed register.
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
bitmore_accumulator() noexcept
|
||||
: acc_(bitmore_detail::zero<Reg>()), bound_(0)
|
||||
{ }
|
||||
|
||||
/// @brief Proven upper bound on every slot in `value()`.
|
||||
HEDLEY_PURE
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
unsigned bound() const noexcept
|
||||
{
|
||||
return bound_;
|
||||
}
|
||||
|
||||
/// @brief Unreduced slots. Each is `≤ bound()` and congruent to the sum.
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
Reg value() const noexcept
|
||||
{
|
||||
return acc_;
|
||||
}
|
||||
|
||||
/// @brief Add one slot.
|
||||
/// @param x addends. Each slot must be `≤ addend_max`.
|
||||
/// @param addend_max worst-case slot in `x`, clamped to `slot_max`. Pass
|
||||
/// `slot_max` when the addend is an arbitrary lane; the accumulator
|
||||
/// partial-reduces it if the slack cannot absorb that.
|
||||
/// \complexity At most four partial reductions and one lane add.
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
void add(Reg x, unsigned addend_max) noexcept
|
||||
{
|
||||
if (addend_max > mod::slot_max)
|
||||
addend_max = mod::slot_max;
|
||||
for (;;)
|
||||
{
|
||||
if (bound_ + addend_max <= mod::slot_max)
|
||||
break;
|
||||
if (bound_ > mod::partial_bound)
|
||||
{
|
||||
acc_ = mod::partial_reduce(acc_);
|
||||
bound_ = mod::partial_bound;
|
||||
continue;
|
||||
}
|
||||
if (addend_max > mod::partial_bound)
|
||||
{
|
||||
x = mod::partial_reduce(x);
|
||||
addend_max = mod::partial_bound;
|
||||
continue;
|
||||
}
|
||||
if (bound_ > mod::stable_bound)
|
||||
{
|
||||
acc_ = mod::partial_reduce(acc_);
|
||||
bound_ = mod::stable_bound;
|
||||
continue;
|
||||
}
|
||||
if (addend_max > mod::stable_bound)
|
||||
{
|
||||
x = mod::partial_reduce(x);
|
||||
addend_max = mod::stable_bound;
|
||||
continue;
|
||||
}
|
||||
break;
|
||||
}
|
||||
if constexpr (LaneBits == 8u)
|
||||
acc_ = bitmore_detail::add_bytes<Reg>(acc_, x);
|
||||
else
|
||||
acc_ = bitmore_detail::add_epi16<Reg>(acc_, x);
|
||||
bound_ += addend_max;
|
||||
}
|
||||
|
||||
/// @brief Fold one BitMore bit: `acc = 2*acc + bit0` inside each slot.
|
||||
/// @param bit bit 0 of each slot is the new low bit. Higher bits are ignored.
|
||||
/// \complexity Up to two partial reductions when the proven bound sets the
|
||||
/// lane's high bit, then a shift and an OR.
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
void insert_bit(Reg bit) noexcept
|
||||
{
|
||||
while (bound_ > (mod::slot_max >> 1))
|
||||
{
|
||||
acc_ = mod::partial_reduce(acc_);
|
||||
bound_ = bound_ > mod::partial_bound ? mod::partial_bound
|
||||
: mod::stable_bound;
|
||||
}
|
||||
if constexpr (LaneBits == 8u)
|
||||
acc_ = bitmore_detail::shift_in_bit<Reg>(acc_, bit);
|
||||
else
|
||||
acc_ = bitmore_detail::shift_in_bit16<Reg>(acc_, bit);
|
||||
bound_ = bound_ * 2u + 1u;
|
||||
}
|
||||
|
||||
/// @brief Full residue of every slot, in `0 .. Modulus-1`.
|
||||
/// \complexity One `full_reduce` of the register.
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_PURE
|
||||
Reg reduced() const noexcept
|
||||
{
|
||||
return mod::full_reduce(acc_);
|
||||
}
|
||||
|
||||
private:
|
||||
Reg acc_;
|
||||
unsigned bound_;
|
||||
};
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_BITMORE_MOD_HPP__
|
||||
|
|
@ -130,7 +130,7 @@ class bitstring : public bit_array_base<bitstring<Nbits, WordT>, WordT>
|
|||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
constexpr explicit bitstring(word_type val) noexcept
|
||||
: data_{Nbits < bits_per_word ? static_cast<word_type>(val & (static_cast<word_type>(~word_type{0}) >> utils::bitlength_of_v<word_type> - Nbits)) : val} { }
|
||||
: data_{Nbits < bits_per_word ? static_cast<word_type>(val & (static_cast<word_type>(~word_type{0}) >> (utils::bitlength_of_v<word_type> - Nbits))) : val} { }
|
||||
|
||||
/// @brief Constructs a `dpf::bitstring` using the characters in the
|
||||
/// `std::basic_string` `str`. An optional starting position `pos`
|
||||
|
|
@ -782,11 +782,22 @@ struct mod_pow_2<dpf::bitstring<Nbits, WordT>>
|
|||
}
|
||||
};
|
||||
|
||||
/// @brief Bitstrings add by XOR over their packed words.
|
||||
template <std::size_t Nbits, typename WordT>
|
||||
struct has_characteristic_two<dpf::bitstring<Nbits, WordT>>
|
||||
: public std::true_type {};
|
||||
|
||||
} // namespace utils
|
||||
|
||||
namespace bitstrings
|
||||
{
|
||||
|
||||
/// @name bitN_t
|
||||
/// @{
|
||||
|
||||
/// @brief `bitN_t` is `dpf::bitstring<N>` for N from 1 through 128. Not `dpf::bit`.
|
||||
/// @see dpf::bitstring
|
||||
/// @see dpf::bit
|
||||
// 1--9
|
||||
using bit1_t = dpf::bitstring<1>;
|
||||
using bit2_t = dpf::bitstring<2>;
|
||||
|
|
@ -929,6 +940,8 @@ using bit126_t = dpf::bitstring<126>;
|
|||
using bit127_t = dpf::bitstring<127>;
|
||||
using bit128_t = dpf::bitstring<128>;
|
||||
|
||||
/// @}
|
||||
|
||||
namespace literals = dpf::literals::bitstrings;
|
||||
|
||||
} // namespace bitstrings
|
||||
|
|
@ -1008,6 +1021,11 @@ constexpr static auto operator "" _bitstring_u128()
|
|||
return bitstring_literal<dpf::bitstring<sizeof...(bits), simde_uint128>, bits...>();
|
||||
}
|
||||
|
||||
/// @name bitstring literals `_bN`
|
||||
/// @{
|
||||
|
||||
/// @brief Digit-string literal for `dpf::bitstring<N>`. The first character is the high bit.
|
||||
/// @see dpf::bitstring
|
||||
// 1--9
|
||||
template <char ...bits> constexpr static auto operator "" _b1() { return bitstring_literal<dpf::bitstring<1>, bits...>(); }
|
||||
template <char ...bits> constexpr static auto operator "" _b2() { return bitstring_literal<dpf::bitstring<2>, bits...>(); }
|
||||
|
|
@ -1150,6 +1168,8 @@ template <char ...bits> constexpr static auto operator "" _b126() { return bitst
|
|||
template <char ...bits> constexpr static auto operator "" _b127() { return bitstring_literal<dpf::bitstring<127>, bits...>(); }
|
||||
template <char ...bits> constexpr static auto operator "" _b128() { return bitstring_literal<dpf::bitstring<128>, bits...>(); }
|
||||
|
||||
/// @}
|
||||
|
||||
} // namespace bitstrings
|
||||
|
||||
} // namespace literals
|
||||
|
|
|
|||
149
include/dpf/blob.hpp
Normal file
149
include/dpf/blob.hpp
Normal file
|
|
@ -0,0 +1,149 @@
|
|||
/// @file dpf/blob.hpp
|
||||
/// @brief Fixed-length XOR byte-string leaf (`dpf::blob<N>`).
|
||||
/// @details A mailbox row or other byte payload that does not fit in one AES
|
||||
/// block. Keygen stretches each final seed with the exterior PRG
|
||||
/// (counter mode) into `N` bytes and XORs a correction of that
|
||||
/// length, matching Express's `genDPF` leaf. The tree and depth are
|
||||
/// unchanged: Boyle packing uses `ceil(8N / λ)` exterior blocks and
|
||||
/// `lg(outputs_per_leaf) = 0`.
|
||||
/// @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_DPF_BLOB_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_BLOB_HPP__
|
||||
|
||||
#include <array>
|
||||
#include <cstddef>
|
||||
#include <cstring>
|
||||
#include <type_traits>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/utils.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
/// @brief XOR share of an `N`-byte string.
|
||||
/// @tparam N byte length (`N >= 1`)
|
||||
template <std::size_t N>
|
||||
struct blob
|
||||
{
|
||||
static_assert(N >= 1, "dpf::blob: N must be at least 1");
|
||||
static constexpr std::size_t size = N;
|
||||
std::array<unsigned char, N> bytes{};
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
constexpr blob() noexcept = default;
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
explicit blob(const unsigned char * src, std::size_t n = N) noexcept
|
||||
{
|
||||
const std::size_t m = n < N ? n : N;
|
||||
if (m != 0)
|
||||
std::memcpy(bytes.data(), src, m);
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
explicit blob(std::array<unsigned char, N> b) noexcept
|
||||
: bytes{std::move(b)}
|
||||
{ }
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
unsigned char * data() noexcept { return bytes.data(); }
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
const unsigned char * data() const noexcept { return bytes.data(); }
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
friend constexpr bool operator==(const blob & a, const blob & b) noexcept
|
||||
{
|
||||
for (std::size_t i = 0; i < N; ++i)
|
||||
if (a.bytes[i] != b.bytes[i])
|
||||
return false;
|
||||
return true;
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
friend constexpr bool operator!=(const blob & a, const blob & b) noexcept
|
||||
{
|
||||
return !(a == b);
|
||||
}
|
||||
|
||||
/// @brief XOR (group addition / subtraction for this leaf).
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
friend blob operator^(const blob & a, const blob & b) noexcept
|
||||
{
|
||||
blob out;
|
||||
for (std::size_t i = 0; i < N; ++i)
|
||||
out.bytes[i] = static_cast<unsigned char>(a.bytes[i] ^ b.bytes[i]);
|
||||
return out;
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
friend blob operator+(const blob & a, const blob & b) noexcept
|
||||
{
|
||||
return a ^ b;
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
friend blob operator-(const blob & a, const blob & b) noexcept
|
||||
{
|
||||
return a ^ b;
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
blob & operator^=(const blob & o) noexcept
|
||||
{
|
||||
*this = *this ^ o;
|
||||
return *this;
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
blob & operator+=(const blob & o) noexcept
|
||||
{
|
||||
return (*this) ^= o;
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
struct is_blob : std::false_type {};
|
||||
template <std::size_t N>
|
||||
struct is_blob<blob<N>> : std::true_type {};
|
||||
template <typename T>
|
||||
inline constexpr bool is_blob_v = is_blob<std::decay_t<T>>::value;
|
||||
|
||||
namespace utils
|
||||
{
|
||||
|
||||
template <std::size_t N>
|
||||
struct bitlength_of<dpf::blob<N>>
|
||||
: public std::integral_constant<std::size_t, N * 8>
|
||||
{ };
|
||||
|
||||
template <std::size_t N, typename NodeT>
|
||||
struct bitlength_of_output<dpf::blob<N>, NodeT>
|
||||
: public std::integral_constant<std::size_t, N * 8>
|
||||
{ };
|
||||
|
||||
template <std::size_t N>
|
||||
struct has_characteristic_two<dpf::blob<N>> : std::true_type
|
||||
{ };
|
||||
|
||||
template <std::size_t N>
|
||||
struct is_xor_wrapper<dpf::blob<N>> : std::true_type
|
||||
{ };
|
||||
|
||||
} // namespace utils
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_BLOB_HPP__
|
||||
|
|
@ -5,6 +5,7 @@
|
|||
/// up to the next checkpoint; a full-domain memoizer already holds
|
||||
/// those nodes. `q` tail bits, when the comparison sets the key
|
||||
/// depth, are a residual table on the node at height `h`.
|
||||
/// @note Boyle, Chandran, Gilboa, Gupta, Ishai, Kumar, and Rathee (EUROCRYPT 2021, ePrint 2020/1392) publish a value-correction word on every level. This implementation is ahead of that DCF on payload size: one ring word every B levels, about B times fewer value words, with the same one-seed-word spine. Point evaluation expands the siblings between checkpoints.
|
||||
/// @copyright Copyright (c) 2019-2026 Ryan Henry and [others](@ref authors)
|
||||
/// @license Released under a GNU General Public v2.0 (GPLv2) license.
|
||||
|
||||
|
|
@ -25,6 +26,7 @@
|
|||
#include "dpf/tree_traits.hpp"
|
||||
#include "dpf/twiddle.hpp"
|
||||
#include "dpf/utils.hpp"
|
||||
#include "dpf/verifiable.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
|
@ -33,6 +35,20 @@ namespace detail
|
|||
namespace blocked
|
||||
{
|
||||
|
||||
/// @brief High bit set on `fold_node` level so blocked proofs diverge from native.
|
||||
/// @details Must match the level passed to `make_cs` at blocked keygen. Parked
|
||||
/// sibling folds reuse this tag with the sibling's tree depth so they
|
||||
/// share `correction_seeds[depth-1]`.
|
||||
inline constexpr std::size_t fold_spine_tag = std::size_t{1} << 15;
|
||||
|
||||
template <typename NodeT>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void fold_spine_node(proof_token & pi, std::size_t level,
|
||||
psnip_uint64_t x_bits, NodeT seed, const cs_block & cs) noexcept
|
||||
{
|
||||
detail::vdpf::fold_node(pi, fold_spine_tag | level, x_bits, seed, cs);
|
||||
}
|
||||
|
||||
template <std::size_t H, std::size_t B>
|
||||
struct schedule
|
||||
{
|
||||
|
|
@ -306,7 +322,8 @@ uint64_t finish_share(const KeyT & dpf, uint64_t suffix, uint64_t acc,
|
|||
}
|
||||
|
||||
template <typename KeyT, typename InputT, typename PathMemoizer>
|
||||
uint64_t eval_share(const KeyT & dpf, InputT tx, PathMemoizer & path)
|
||||
uint64_t eval_share(const KeyT & dpf, InputT tx, PathMemoizer & path,
|
||||
proof_token * pi = nullptr)
|
||||
{
|
||||
using node = typename KeyT::interior_node;
|
||||
const auto & ch = dpf.cmp();
|
||||
|
|
@ -326,18 +343,42 @@ uint64_t eval_share(const KeyT & dpf, InputT tx, PathMemoizer & path)
|
|||
const std::size_t nbits = static_cast<std::size_t>(ch.nbits);
|
||||
const int party = dpf::get_lo_bit(dpf.root()) ? 1 : 0;
|
||||
|
||||
// Expand the seed spine without native path folds; blocked proofs use
|
||||
// domain-separated tags on newly filled levels only (path-memo safe).
|
||||
const auto resume = dpf::detail::path_resume_for_level(path, dpf, tx, h);
|
||||
dpf::detail::ensure_level(dpf, tx, path, h);
|
||||
|
||||
if constexpr (KeyT::is_verifiable)
|
||||
{
|
||||
if (pi != nullptr && resume <= h)
|
||||
{
|
||||
constexpr auto input_bits =
|
||||
utils::bitlength_of_v<typename KeyT::input_type>;
|
||||
for (std::size_t level = resume; level <= h; ++level)
|
||||
{
|
||||
const auto x_bits = static_cast<psnip_uint64_t>(
|
||||
utils::to_integral_type<typename KeyT::input_type>{}(tx)
|
||||
>> (input_bits - level));
|
||||
fold_spine_node(*pi, level - 1, x_bits, path[level],
|
||||
dpf.correction_seeds()[level - 1]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct parked
|
||||
{
|
||||
node seed;
|
||||
std::size_t depth;
|
||||
psnip_uint64_t prefix_bits;
|
||||
};
|
||||
parked pend[128];
|
||||
std::size_t npend = 0;
|
||||
|
||||
uint64_t acc = 0;
|
||||
auto bit_mask = KeyT::msb_mask;
|
||||
// Value accumulation always walks from the root. Proof folds for parked
|
||||
// siblings must not repeat depths already authenticated on a warm path:
|
||||
// re-folding XORs the same contribution away (path-memo cancel bug).
|
||||
for (std::size_t level = 0; level < h; ++level, bit_mask >>= 1)
|
||||
{
|
||||
const bool xi = !!(bit_mask & tx);
|
||||
|
|
@ -349,6 +390,13 @@ uint64_t eval_share(const KeyT & dpf, InputT tx, PathMemoizer & path)
|
|||
{
|
||||
pend[npend].seed = right;
|
||||
pend[npend].depth = level + 1;
|
||||
// Sibling is the right child of `parent`: path bits with low bit 1.
|
||||
const auto path_bits = static_cast<psnip_uint64_t>(
|
||||
utils::to_integral_type<typename KeyT::input_type>{}(tx)
|
||||
>> (utils::bitlength_of_v<typename KeyT::input_type>
|
||||
- (level + 1)));
|
||||
pend[npend].prefix_bits = (path_bits & ~static_cast<psnip_uint64_t>(1))
|
||||
| static_cast<psnip_uint64_t>(1);
|
||||
++npend;
|
||||
}
|
||||
const std::size_t c = level + 1;
|
||||
|
|
@ -357,6 +405,15 @@ uint64_t eval_share(const KeyT & dpf, InputT tx, PathMemoizer & path)
|
|||
const uint64_t word = dpf.value_cw(sched::index(c));
|
||||
for (std::size_t p = 0; p < npend; ++p)
|
||||
{
|
||||
if constexpr (KeyT::is_verifiable)
|
||||
{
|
||||
if (pi != nullptr && pend[p].depth >= resume)
|
||||
{
|
||||
// Tree depth (not checkpoint index): must match CS level.
|
||||
fold_spine_node(*pi, pend[p].depth - 1, pend[p].prefix_bits,
|
||||
pend[p].seed, dpf.correction_seeds()[pend[p].depth - 1]);
|
||||
}
|
||||
}
|
||||
acc = add_frontier<KeyT>(acc, pend[p].seed, pend[p].depth, c,
|
||||
dpf, word, mask, party);
|
||||
}
|
||||
|
|
@ -392,9 +449,17 @@ const typename KeyT::interior_node & memo_node(const Memo & memo, Integral prefi
|
|||
return memo[depth][idx];
|
||||
}
|
||||
|
||||
template <typename KeyT, typename Integral, typename Memo>
|
||||
/// @brief Evaluate one lane from an interval memo; optionally fold a VDPF proof.
|
||||
/// @details When `pi == nullptr` or the key is not verifiable, matches the
|
||||
/// historical body (no folds). When set, folds path / covered
|
||||
/// checkpoint / frontier seeds with `fold_spine_node`, skipping depths
|
||||
/// already covered by `path` (same resume rule as `eval_share`).
|
||||
/// Interval BFS prove must keep passing a null `pi` here.
|
||||
template <typename KeyT, typename Integral, typename Memo,
|
||||
typename PathMemoizer = ::dpf::basic_path_memoizer<KeyT>>
|
||||
uint64_t eval_share_memo(const KeyT & dpf, Integral lane,
|
||||
Integral from_lane, Integral to_excl, const Memo & memo)
|
||||
Integral from_lane, Integral to_excl, const Memo & memo,
|
||||
proof_token * pi = nullptr, PathMemoizer * path = nullptr)
|
||||
{
|
||||
using node = typename KeyT::interior_node;
|
||||
const auto & ch = dpf.cmp();
|
||||
|
|
@ -412,6 +477,43 @@ uint64_t eval_share_memo(const KeyT & dpf, Integral lane,
|
|||
constexpr std::size_t h = KeyT::cmp_h;
|
||||
using sched = schedule<h, KeyT::cmp_block>;
|
||||
const int party = dpf::get_lo_bit(dpf.root()) ? 1 : 0;
|
||||
const std::size_t nbits = static_cast<std::size_t>(ch.nbits);
|
||||
|
||||
std::size_t resume = 0;
|
||||
if constexpr (KeyT::is_verifiable)
|
||||
{
|
||||
if (pi != nullptr && path != nullptr)
|
||||
{
|
||||
const auto tx = static_cast<typename KeyT::input_type>(lane);
|
||||
resume = dpf::detail::path_resume_for_level(*path, dpf, tx, h);
|
||||
dpf::detail::path_note_filled_to(*path, h);
|
||||
}
|
||||
}
|
||||
|
||||
if constexpr (KeyT::is_verifiable)
|
||||
{
|
||||
if (pi != nullptr && resume <= h)
|
||||
{
|
||||
Integral pfx = 0;
|
||||
for (std::size_t depth = 0; depth <= h; ++depth)
|
||||
{
|
||||
if (depth >= resume && depth >= 1)
|
||||
{
|
||||
fold_spine_node(*pi, depth - 1,
|
||||
static_cast<psnip_uint64_t>(pfx),
|
||||
memo_node<KeyT>(memo, pfx, depth, from_lane),
|
||||
dpf.correction_seeds()[depth - 1]);
|
||||
}
|
||||
if (depth < h)
|
||||
{
|
||||
const bool xi =
|
||||
((lane >> (nbits - 1 - depth)) & Integral{1}) != 0;
|
||||
pfx = static_cast<Integral>(
|
||||
(pfx << 1) | (xi ? Integral{1} : Integral{0}));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct parked
|
||||
{
|
||||
|
|
@ -424,7 +526,6 @@ uint64_t eval_share_memo(const KeyT & dpf, Integral lane,
|
|||
|
||||
uint64_t acc = 0;
|
||||
Integral path_pref = 0;
|
||||
const std::size_t nbits = static_cast<std::size_t>(ch.nbits);
|
||||
for (std::size_t level = 0; level < h; ++level)
|
||||
{
|
||||
const bool xi = ((lane >> (nbits - 1 - level)) & Integral{1}) != 0;
|
||||
|
|
@ -440,7 +541,8 @@ uint64_t eval_share_memo(const KeyT & dpf, Integral lane,
|
|||
pend[npend].depth = level + 1;
|
||||
++npend;
|
||||
}
|
||||
path_pref = static_cast<Integral>((path_pref << 1) | Integral{xi ? 1 : 0});
|
||||
path_pref = static_cast<Integral>(
|
||||
(path_pref << 1) | (xi ? Integral{1} : Integral{0}));
|
||||
const std::size_t c = level + 1;
|
||||
if (!sched::contains(c))
|
||||
continue;
|
||||
|
|
@ -461,6 +563,16 @@ uint64_t eval_share_memo(const KeyT & dpf, Integral lane,
|
|||
for (Integral k = 0; k < nleaf; ++k)
|
||||
{
|
||||
const auto pref = static_cast<Integral>(leftmost + k);
|
||||
if constexpr (KeyT::is_verifiable)
|
||||
{
|
||||
if (pi != nullptr && c >= resume && c >= 1)
|
||||
{
|
||||
fold_spine_node(*pi, c - 1,
|
||||
static_cast<psnip_uint64_t>(pref),
|
||||
memo_node<KeyT>(memo, pref, c, from_lane),
|
||||
dpf.correction_seeds()[c - 1]);
|
||||
}
|
||||
}
|
||||
acc = add_membership<KeyT>(acc,
|
||||
memo_node<KeyT>(memo, pref, c, from_lane), word, mask,
|
||||
party);
|
||||
|
|
@ -468,6 +580,17 @@ uint64_t eval_share_memo(const KeyT & dpf, Integral lane,
|
|||
}
|
||||
else
|
||||
{
|
||||
if constexpr (KeyT::is_verifiable)
|
||||
{
|
||||
if (pi != nullptr && pend[p].depth >= resume
|
||||
&& pend[p].depth >= 1)
|
||||
{
|
||||
fold_spine_node(*pi, pend[p].depth - 1,
|
||||
static_cast<psnip_uint64_t>(pend[p].prefix),
|
||||
pend[p].seed,
|
||||
dpf.correction_seeds()[pend[p].depth - 1]);
|
||||
}
|
||||
}
|
||||
acc = add_frontier<KeyT>(acc, pend[p].seed, pend[p].depth, c,
|
||||
dpf, word, mask, party);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@
|
|||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/aligned_allocator.hpp"
|
||||
#include "dpf/experiment_note.hpp"
|
||||
#include "dpf/prg.hpp"
|
||||
#include "dpf/random.hpp"
|
||||
|
||||
|
|
@ -50,6 +51,19 @@ typename PRG::block_type mask_master(typename PRG::block_type master) noexcept
|
|||
return out;
|
||||
}
|
||||
|
||||
/// @brief Master for IT-MAC tag masks. Distinct from `mask_master`.
|
||||
template <typename PRG>
|
||||
HEDLEY_NO_THROW
|
||||
typename PRG::block_type tag_master(typename PRG::block_type master) noexcept
|
||||
{
|
||||
unsigned char raw[sizeof(master)];
|
||||
std::memcpy(raw, &master, sizeof(master));
|
||||
raw[sizeof(master) - 1] ^= 0x02u;
|
||||
typename PRG::block_type out{};
|
||||
std::memcpy(&out, raw, sizeof(out));
|
||||
return out;
|
||||
}
|
||||
|
||||
template <typename PRG, typename T>
|
||||
struct lane_codec
|
||||
{
|
||||
|
|
@ -192,12 +206,16 @@ public:
|
|||
explicit buffered_prg(std::size_t per_stream_buffer_elems = 1024u)
|
||||
: seed_(sample_master_seed<PRG>()),
|
||||
buffers_(make_buffers(per_stream_buffer_elems))
|
||||
{ }
|
||||
{
|
||||
note_experiment_seed("buffered_prg", seed_);
|
||||
}
|
||||
|
||||
explicit buffered_prg(seed_type seed, std::size_t per_stream_buffer_elems = 1024u)
|
||||
: seed_(seed),
|
||||
buffers_(make_buffers(per_stream_buffer_elems))
|
||||
{ }
|
||||
{
|
||||
note_experiment_seed("buffered_prg", seed_);
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
const seed_type & seed() const noexcept { return seed_; }
|
||||
|
|
@ -282,6 +300,7 @@ public:
|
|||
{
|
||||
if (window_ == 0)
|
||||
throw std::invalid_argument("prg lane window must be positive");
|
||||
note_experiment_seed("lane_table", seed_);
|
||||
}
|
||||
|
||||
lane_table(const lane_table &) = delete;
|
||||
|
|
|
|||
129
include/dpf/caller_fold.hpp
Normal file
129
include/dpf/caller_fold.hpp
Normal file
|
|
@ -0,0 +1,129 @@
|
|||
/// @file dpf/caller_fold.hpp
|
||||
/// @brief One-pass caller fold for the full-domain / point / sequence walks.
|
||||
/// @details `sketch_ref` folds every extractable share into three `fp61`
|
||||
/// moments. A protocol whose audit lives in a different group (the
|
||||
/// Express multiplication proof, Pika's Schwartz–Zippel check, a
|
||||
/// keyword-PIR record XOR) wants the same one-pass hook without a
|
||||
/// second expansion and without the `fp61` conversion. These helpers
|
||||
/// take an inlined `fold(index, share)` callable, invoked once per
|
||||
/// written output in the same loop that writes the buffer. The fold is
|
||||
/// a template argument, not a virtual call, so it inlines exactly the
|
||||
/// way `sketch_ref::absorb` does today. `dpf::sketch` stays the
|
||||
/// weight-1 `fp61` fold; existing call sites are untouched.
|
||||
/// @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_DPF_CALLER_FOLD_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_CALLER_FOLD_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <iterator>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/dpf_key.hpp"
|
||||
#include "dpf/eval_common.hpp"
|
||||
#include "dpf/eval_full.hpp"
|
||||
#include "dpf/eval_point.hpp"
|
||||
#include "dpf/eval_target.hpp"
|
||||
#include "dpf/eval_walk.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
namespace detail_fold
|
||||
{
|
||||
|
||||
/// @brief True when `Fold` is callable as `fold(std::size_t, share)` for the
|
||||
/// group element written by output `I` of `DpfKey`.
|
||||
template <typename Fold, typename DpfKey, std::size_t I, typename = void>
|
||||
struct is_index_fold : std::false_type {};
|
||||
|
||||
template <typename Fold, typename DpfKey, std::size_t I>
|
||||
struct is_index_fold<Fold, DpfKey, I,
|
||||
std::void_t<decltype(std::declval<Fold &>()(
|
||||
std::declval<std::size_t>(),
|
||||
std::declval<typename DpfKey::template concrete_output_type<I>>()))>>
|
||||
: std::true_type {};
|
||||
|
||||
template <typename Fold, typename DpfKey, std::size_t I>
|
||||
inline constexpr bool is_index_fold_v =
|
||||
is_index_fold<std::decay_t<Fold>, DpfKey, I>::value;
|
||||
|
||||
} // namespace detail_fold
|
||||
|
||||
// N.B.: the one-pass `eval_full_add_into(buf, key, fold)` hook lives in
|
||||
// `dpf/eval_walk.hpp` alongside the `rotate` / `sketch_ref` overloads. This
|
||||
// header adds the remaining fold-carrying walks (allocate-and-fold full eval,
|
||||
// point eval, and the full-domain keyword-PIR XOR) so every walk that accepts
|
||||
// a `sketch_ref` also accepts a generic caller fold.
|
||||
|
||||
/// @brief Full-domain expansion of output `I`, folding each written share.
|
||||
/// @details Allocates a fresh buffer (like `eval_full(key)`), folds every point
|
||||
/// into `fold`, and returns the `(buffer, iterable)` pair so the caller
|
||||
/// keeps the expanded shares as well as the audit.
|
||||
/// \complexity One full-domain expansion, `Θ(2^n)`, plus one fold per point.
|
||||
template <std::size_t I = 0, typename DpfKey, typename Fold,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>
|
||||
&& detail_fold::is_index_fold_v<Fold, DpfKey, I>, bool> = true>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_full_fold(const DpfKey & dpf, Fold fold)
|
||||
{
|
||||
auto result = eval_full<I>(dpf);
|
||||
auto & iter = result.second;
|
||||
std::size_t i = 0;
|
||||
for (auto it = std::begin(iter); it != std::end(iter); ++it, ++i)
|
||||
fold(i, detail_walk::group_value(*it));
|
||||
return result;
|
||||
}
|
||||
|
||||
/// @brief Evaluate output `I` at `x` and fold the single written share.
|
||||
/// @details The fold is called once, with index `0`, matching the one output
|
||||
/// `eval_point` writes.
|
||||
/// \complexity O(n) time; one interior traversal per level.
|
||||
template <std::size_t I = 0, typename DpfKey, typename InputT, typename Fold,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>
|
||||
&& detail_fold::is_index_fold_v<Fold, DpfKey, I>, bool> = true>
|
||||
auto eval_point_fold(const DpfKey & dpf, InputT && x, Fold fold)
|
||||
{
|
||||
auto out = eval_point<I>(dpf, std::forward<InputT>(x));
|
||||
fold(std::size_t{0}, detail_walk::group_value(*out));
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief Keyword-PIR XOR fold: `⊕ records[i]` over `i` where `DPF_I(i)` is set.
|
||||
/// @details Runs one full-domain bit expansion of output `I` and, in that same
|
||||
/// loop, XORs `records[i]` into a running accumulator whenever the
|
||||
/// party's bit share at `i` is 1 — without ever materializing the bit
|
||||
/// vector. Each server returns its share of the XOR; the two shares
|
||||
/// reconstruct (XOR) to `⊕ records[i]` over the 1-set of the DPF, i.e.
|
||||
/// the matched record for a point key. `records` is indexed by domain
|
||||
/// point and must cover the domain (or at least the evaluated prefix).
|
||||
/// \complexity One full-domain expansion, `Θ(2^n)`, plus one XOR per set bit.
|
||||
template <std::size_t I = 0, typename DpfKey, typename Records,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_sequence_xor(const DpfKey & key, const Records & records)
|
||||
{
|
||||
using record_type = std::decay_t<decltype(records[std::size_t{0}])>;
|
||||
auto result = eval_full<I>(key);
|
||||
auto & iter = result.second;
|
||||
record_type acc{};
|
||||
std::size_t i = 0;
|
||||
const std::size_t n = static_cast<std::size_t>(std::size(records));
|
||||
for (auto it = std::begin(iter); it != std::end(iter) && i < n; ++it, ++i)
|
||||
{
|
||||
if (static_cast<bool>(*it))
|
||||
acc = static_cast<record_type>(acc ^ records[i]);
|
||||
}
|
||||
return acc;
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_CALLER_FOLD_HPP__
|
||||
552
include/dpf/circuit.hpp
Normal file
552
include/dpf/circuit.hpp
Normal file
|
|
@ -0,0 +1,552 @@
|
|||
/// @file dpf/circuit.hpp
|
||||
/// @brief One arithmetic circuit. Each party binds its shares and drives.
|
||||
/// @details Prep is a `prep::cursor` over a view from `deal_views`, a file, a
|
||||
/// `stream_array`, or `setup_2pc_sampled`. Compare, truncation, mux,
|
||||
/// and GMW A2B are instructions whose prep atoms come from that view.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_CIRCUIT_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_CIRCUIT_HPP__
|
||||
|
||||
#include <chrono>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <stdexcept>
|
||||
#include <vector>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/edabit.hpp"
|
||||
#include "dpf/net/round_sink.hpp"
|
||||
#include "dpf/net/stream_array.hpp"
|
||||
#include "dpf/prep_source.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace mpc
|
||||
{
|
||||
|
||||
struct wire
|
||||
{
|
||||
std::uint32_t id = 0;
|
||||
};
|
||||
|
||||
/// @brief Recorded program plus the prep it will consume, in order.
|
||||
class circuit
|
||||
{
|
||||
public:
|
||||
explicit circuit(std::uint16_t limb = 8)
|
||||
: limb_(limb)
|
||||
{
|
||||
if (limb_ == 0 || limb_ > 8)
|
||||
throw std::invalid_argument("circuit limb must be 1..8");
|
||||
dem_.limb = limb_;
|
||||
}
|
||||
|
||||
std::uint16_t limb() const noexcept { return limb_; }
|
||||
const prep::demand & prep() const noexcept { return dem_; }
|
||||
|
||||
/// @brief One slot per online exchange, in program order.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::vector<std::size_t> slot_bytes() const { return slots_; }
|
||||
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
wire input() { return alloc(op::in); }
|
||||
|
||||
/// @brief Owner holds the secret and sends a fresh mask; the peer stores it.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
wire priv_input(unsigned owner)
|
||||
{
|
||||
auto w = alloc(op::priv_in);
|
||||
code_.back().aux = static_cast<std::uint16_t>(owner);
|
||||
slots_.push_back(limb_);
|
||||
return w;
|
||||
}
|
||||
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
wire add(wire a, wire b) { return bin(op::add, a, b); }
|
||||
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
wire mul(wire a, wire b)
|
||||
{
|
||||
dem_.ring_triples++;
|
||||
auto w = bin(op::mul, a, b);
|
||||
slots_.push_back(limb_);
|
||||
slots_.push_back(limb_);
|
||||
return w;
|
||||
}
|
||||
|
||||
/// @brief GMW AND on XOR bit shares. One bit triple; one open of two masks.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
wire gmw_and(wire p, wire q)
|
||||
{
|
||||
dem_.bit_triples++;
|
||||
auto w = bin(op::gmw_and, p, q);
|
||||
slots_.push_back(2); // (d,e) byte pair
|
||||
return w;
|
||||
}
|
||||
|
||||
/// @brief A2B via GMW carry. `width` bits; opens `x-r` once then AND chain.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
wire a2b(wire x, unsigned width)
|
||||
{
|
||||
if (width == 0 || width > 64)
|
||||
throw std::invalid_argument("circuit a2b width");
|
||||
dem_.dabits += width; // edaBit bits consumed as dabits budget proxy
|
||||
dem_.bit_triples += width;
|
||||
auto w = alloc(op::a2b);
|
||||
code_.back().a = x.id;
|
||||
code_.back().aux = static_cast<std::uint16_t>(width);
|
||||
slots_.push_back(limb_); // open x - r
|
||||
for (unsigned i = 0; i < width; ++i)
|
||||
slots_.push_back(2); // AND masks
|
||||
return w;
|
||||
}
|
||||
|
||||
/// @brief Reserve prep and opens for an unsigned compare.
|
||||
/// @details The online walk consumes those slots. It does not yet return
|
||||
/// the predicate. Use `share_cmp` or a DCF key for a real compare.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
wire gt(wire x, wire y, unsigned width)
|
||||
{
|
||||
if (width == 0 || width > 64)
|
||||
throw std::invalid_argument("circuit gt width");
|
||||
dem_.dabits += 2u * width;
|
||||
dem_.bit_triples += 3u * width; // a2b carries + compare ANDs (bound)
|
||||
auto w = bin(op::gt, x, y);
|
||||
code_.back().aux = static_cast<std::uint16_t>(width);
|
||||
slots_.push_back(limb_);
|
||||
slots_.push_back(limb_);
|
||||
for (unsigned i = 0; i < 3u * width; ++i)
|
||||
slots_.push_back(2);
|
||||
return w;
|
||||
}
|
||||
|
||||
/// @brief Exact trunc: open `x-r`, then GMW wrap bit.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
wire trunc_exact(wire x, unsigned n, unsigned s)
|
||||
{
|
||||
if (s >= n || n > 64)
|
||||
throw std::invalid_argument("circuit trunc_exact");
|
||||
dem_.dabits += s;
|
||||
dem_.bit_triples += s;
|
||||
auto w = alloc(op::trunc_exact);
|
||||
code_.back().a = x.id;
|
||||
code_.back().aux = static_cast<std::uint16_t>(n);
|
||||
code_.back().aux2 = static_cast<std::uint16_t>(s);
|
||||
slots_.push_back(limb_);
|
||||
for (unsigned i = 0; i < s; ++i)
|
||||
slots_.push_back(2);
|
||||
return w;
|
||||
}
|
||||
|
||||
/// @brief Mux `b + sel·(a-b)` via one bit×ring inject (ring triple + open).
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
wire mux(wire sel, wire a, wire b)
|
||||
{
|
||||
dem_.ring_triples++;
|
||||
dem_.bit_triples++; // bit×ring uses a bit share of sel
|
||||
auto w = alloc(op::mux);
|
||||
code_.back().a = sel.id;
|
||||
code_.back().b = a.id;
|
||||
code_.back().aux = static_cast<std::uint16_t>(b.id);
|
||||
slots_.push_back(limb_);
|
||||
slots_.push_back(limb_);
|
||||
return w;
|
||||
}
|
||||
|
||||
/// @brief Declassify. Two parties exchange and add. Three-party open is
|
||||
/// `run::declassify_ring` / `run::declassify_streams`.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
wire open(wire a)
|
||||
{
|
||||
auto w = alloc(op::open);
|
||||
code_.back().a = a.id;
|
||||
slots_.push_back(limb_);
|
||||
return w;
|
||||
}
|
||||
|
||||
std::uint32_t wire_count() const noexcept { return nwire_; }
|
||||
|
||||
enum class op : std::uint8_t
|
||||
{
|
||||
in,
|
||||
priv_in,
|
||||
add,
|
||||
mul,
|
||||
open,
|
||||
gmw_and,
|
||||
a2b,
|
||||
gt,
|
||||
trunc_exact,
|
||||
mux
|
||||
};
|
||||
|
||||
struct inst
|
||||
{
|
||||
op code = op::in;
|
||||
std::uint32_t dst = 0;
|
||||
std::uint32_t a = 0;
|
||||
std::uint32_t b = 0;
|
||||
std::uint16_t aux = 0;
|
||||
std::uint16_t aux2 = 0;
|
||||
};
|
||||
|
||||
const std::vector<inst> & program() const noexcept { return code_; }
|
||||
|
||||
private:
|
||||
wire alloc(op code)
|
||||
{
|
||||
inst in;
|
||||
in.code = code;
|
||||
in.dst = nwire_++;
|
||||
code_.push_back(in);
|
||||
return wire{in.dst};
|
||||
}
|
||||
|
||||
wire bin(op code, wire a, wire b)
|
||||
{
|
||||
auto w = alloc(code);
|
||||
code_.back().a = a.id;
|
||||
code_.back().b = b.id;
|
||||
return w;
|
||||
}
|
||||
|
||||
std::uint16_t limb_ = 8;
|
||||
std::uint32_t nwire_ = 0;
|
||||
prep::demand dem_{};
|
||||
std::vector<inst> code_;
|
||||
std::vector<std::size_t> slots_;
|
||||
};
|
||||
|
||||
/// @brief One party's evaluation of a `circuit`.
|
||||
class party
|
||||
{
|
||||
public:
|
||||
party(const circuit & c, unsigned id)
|
||||
: circ_(&c),
|
||||
id_(id),
|
||||
limb_(c.limb()),
|
||||
mask_(limb_ >= 8 ? ~std::uint64_t{0}
|
||||
: ((std::uint64_t{1} << (8u * limb_)) - 1u)),
|
||||
words_(c.wire_count(), 0)
|
||||
{
|
||||
if (id_ > 2)
|
||||
throw std::invalid_argument("party id");
|
||||
}
|
||||
|
||||
void bind(wire w, std::uint64_t share)
|
||||
{
|
||||
words_.at(w.id) = share & mask_;
|
||||
}
|
||||
|
||||
void bind_priv(wire w, std::uint64_t secret)
|
||||
{
|
||||
priv_.resize(circ_->wire_count());
|
||||
priv_set_.resize(circ_->wire_count());
|
||||
priv_.at(w.id) = secret & mask_;
|
||||
priv_set_.at(w.id) = 1;
|
||||
}
|
||||
|
||||
void run(net::RoundSink & sink, prep::cursor prep)
|
||||
{
|
||||
if (prep.limb() != limb_)
|
||||
throw std::invalid_argument("prep limb does not match circuit");
|
||||
std::uint16_t round = 0;
|
||||
for (const auto & in : circ_->program())
|
||||
{
|
||||
switch (in.code)
|
||||
{
|
||||
case circuit::op::in:
|
||||
break;
|
||||
case circuit::op::priv_in:
|
||||
words_[in.dst] = priv_in(in, sink, round);
|
||||
break;
|
||||
case circuit::op::add:
|
||||
words_[in.dst] = (words_[in.a] + words_[in.b]) & mask_;
|
||||
break;
|
||||
case circuit::op::mul:
|
||||
words_[in.dst] = mul(in, sink, prep, round);
|
||||
break;
|
||||
case circuit::op::open:
|
||||
words_[in.dst] = exch_sum(words_[in.a], sink, round);
|
||||
break;
|
||||
case circuit::op::gmw_and:
|
||||
words_[in.dst] = gmw_and(in, sink, prep, round);
|
||||
break;
|
||||
case circuit::op::a2b:
|
||||
words_[in.dst] = a2b_op(in, sink, prep, round);
|
||||
break;
|
||||
case circuit::op::gt:
|
||||
words_[in.dst] = gt_op(in, sink, prep, round);
|
||||
break;
|
||||
case circuit::op::trunc_exact:
|
||||
words_[in.dst] = trunc_exact_op(in, sink, prep, round);
|
||||
break;
|
||||
case circuit::op::mux:
|
||||
words_[in.dst] = mux_op(in, sink, prep, round);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Read a packed prep blob from `dealer[0]` then run.
|
||||
void run(net::RoundSink & sink, net::stream_array & dealer,
|
||||
std::size_t prep_bytes)
|
||||
{
|
||||
run(sink, prep::cursor_from_stream(dealer, 0, prep_bytes));
|
||||
}
|
||||
|
||||
std::uint64_t read(wire w) const { return words_.at(w.id); }
|
||||
|
||||
/// @brief How long one round may wait for the peer's bytes.
|
||||
void set_wait_timeout(std::chrono::milliseconds budget) { wait_ = budget; }
|
||||
|
||||
private:
|
||||
std::uint64_t exch_sum(std::uint64_t mine, net::RoundSink & sink,
|
||||
std::uint16_t & round)
|
||||
{
|
||||
std::uint8_t buf[8]{};
|
||||
std::memcpy(buf, &mine, limb_);
|
||||
sink.submit(round, 0, buf, limb_);
|
||||
sink.flush();
|
||||
wait_peer(sink, round);
|
||||
std::uint8_t peer[8]{};
|
||||
sink.read_peer(round, 0, peer, limb_);
|
||||
++round;
|
||||
std::uint64_t o = 0;
|
||||
std::memcpy(&o, peer, limb_);
|
||||
return (mine + o) & mask_;
|
||||
}
|
||||
|
||||
void wait_peer(net::RoundSink & sink, std::uint16_t round) const
|
||||
{
|
||||
net::wait_peer_ready(sink, round, 0, wait_, "circuit");
|
||||
}
|
||||
|
||||
std::uint64_t priv_in(const circuit::inst & in, net::RoundSink & sink,
|
||||
std::uint16_t & round)
|
||||
{
|
||||
if (id_ == in.aux)
|
||||
{
|
||||
if (in.dst >= priv_set_.size() || !priv_set_[in.dst])
|
||||
throw std::logic_error("circuit: owner did not bind_priv");
|
||||
const auto r = dpf::uniform_sample<std::uint64_t>() & mask_;
|
||||
std::uint8_t buf[8]{};
|
||||
std::memcpy(buf, &r, limb_);
|
||||
sink.submit(round, 0, buf, limb_);
|
||||
sink.flush();
|
||||
wait_peer(sink, round);
|
||||
std::uint8_t ignore[8]{};
|
||||
sink.read_peer(round, 0, ignore, limb_);
|
||||
++round;
|
||||
return (priv_[in.dst] - r) & mask_;
|
||||
}
|
||||
std::uint8_t zeros[8]{};
|
||||
sink.submit(round, 0, zeros, limb_);
|
||||
sink.flush();
|
||||
wait_peer(sink, round);
|
||||
std::uint8_t peer[8]{};
|
||||
sink.read_peer(round, 0, peer, limb_);
|
||||
++round;
|
||||
std::uint64_t r = 0;
|
||||
std::memcpy(&r, peer, limb_);
|
||||
return r & mask_;
|
||||
}
|
||||
|
||||
std::uint64_t mul(const circuit::inst & in, net::RoundSink & sink,
|
||||
prep::cursor & prep, std::uint16_t & round)
|
||||
{
|
||||
std::uint8_t ab[8]{}, bb[8]{}, cb[8]{};
|
||||
prep.take_ring(ab, bb, cb);
|
||||
std::uint64_t a = 0, b = 0, c = 0;
|
||||
std::memcpy(&a, ab, limb_);
|
||||
std::memcpy(&b, bb, limb_);
|
||||
std::memcpy(&c, cb, limb_);
|
||||
const auto d = exch_sum((words_[in.a] - a) & mask_, sink, round);
|
||||
const auto e = exch_sum((words_[in.b] - b) & mask_, sink, round);
|
||||
std::uint64_t z = (c + d * b + e * a) & mask_;
|
||||
if (id_ == 0)
|
||||
z = (z + d * e) & mask_;
|
||||
return z;
|
||||
}
|
||||
|
||||
std::uint64_t gmw_and(const circuit::inst & in, net::RoundSink & sink,
|
||||
prep::cursor & prep, std::uint16_t & round)
|
||||
{
|
||||
auto t = prep.take_bit();
|
||||
const std::uint8_t p = static_cast<std::uint8_t>(words_[in.a] & 1u);
|
||||
const std::uint8_t q = static_cast<std::uint8_t>(words_[in.b] & 1u);
|
||||
std::uint8_t mine[2] = {
|
||||
static_cast<std::uint8_t>((p ^ t.a) & 1u),
|
||||
static_cast<std::uint8_t>((q ^ t.b) & 1u)};
|
||||
sink.submit(round, 0, mine, 2);
|
||||
sink.flush();
|
||||
wait_peer(sink, round);
|
||||
std::uint8_t peer[2]{};
|
||||
sink.read_peer(round, 0, peer, 2);
|
||||
++round;
|
||||
const std::uint8_t d = static_cast<std::uint8_t>(mine[0] ^ peer[0]);
|
||||
const std::uint8_t e = static_cast<std::uint8_t>(mine[1] ^ peer[1]);
|
||||
return edabit::and_finish(t, d, e, id_);
|
||||
}
|
||||
|
||||
/// @brief A2B: open `x - r`, then GMW carry chain (one AND per bit).
|
||||
std::uint64_t a2b_op(const circuit::inst & in, net::RoundSink & sink,
|
||||
prep::cursor & prep, std::uint16_t & round)
|
||||
{
|
||||
const unsigned width = in.aux;
|
||||
std::uint64_t r_arith = 0;
|
||||
std::vector<std::uint8_t> r_bits(width);
|
||||
for (unsigned i = 0; i < width; ++i)
|
||||
{
|
||||
auto d = prep.take_dabit();
|
||||
r_bits[i] = static_cast<std::uint8_t>(d.bit & 1u);
|
||||
r_arith = (r_arith + ((d.arith & 1u) << i)) & mask_;
|
||||
}
|
||||
const auto delta =
|
||||
exch_sum((words_[in.a] - r_arith) & mask_, sink, round);
|
||||
std::uint8_t carry = 0;
|
||||
std::uint64_t out = 0;
|
||||
for (unsigned i = 0; i < width; ++i)
|
||||
{
|
||||
auto t = prep.take_bit();
|
||||
const std::uint8_t di =
|
||||
static_cast<std::uint8_t>((delta >> i) & 1u);
|
||||
const std::uint8_t ri = r_bits[i];
|
||||
std::uint8_t mine[2] = {
|
||||
static_cast<std::uint8_t>((ri ^ t.a) & 1u),
|
||||
static_cast<std::uint8_t>((carry ^ t.b) & 1u)};
|
||||
sink.submit(round, 0, mine, 2);
|
||||
sink.flush();
|
||||
wait_peer(sink, round);
|
||||
std::uint8_t peer[2]{};
|
||||
sink.read_peer(round, 0, peer, 2);
|
||||
++round;
|
||||
const std::uint8_t d = static_cast<std::uint8_t>(mine[0] ^ peer[0]);
|
||||
const std::uint8_t e = static_cast<std::uint8_t>(mine[1] ^ peer[1]);
|
||||
const std::uint8_t rc = edabit::and_finish(t, d, e, id_);
|
||||
const std::uint8_t bi = static_cast<std::uint8_t>(
|
||||
((id_ == 0 ? (ri ^ di ^ carry) : (ri ^ carry))) & 1u);
|
||||
const std::uint8_t rd = static_cast<std::uint8_t>(di & ri);
|
||||
const std::uint8_t dc = static_cast<std::uint8_t>(di & carry);
|
||||
carry = static_cast<std::uint8_t>((rd ^ dc ^ rc) & 1u);
|
||||
out |= (static_cast<std::uint64_t>(bi) << i);
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
std::uint64_t gt_op(const circuit::inst & in, net::RoundSink & sink,
|
||||
prep::cursor & prep, std::uint16_t & round)
|
||||
{
|
||||
// Recorded as a sequence of opens + AND rounds; evaluate via two A2B
|
||||
// slot walks then a bit compare consuming remaining triples.
|
||||
circuit::inst ax{circuit::op::a2b, 0, in.a, 0, in.aux, 0};
|
||||
circuit::inst ay{circuit::op::a2b, 0, in.b, 0, in.aux, 0};
|
||||
(void)a2b_op(ax, sink, prep, round);
|
||||
(void)a2b_op(ay, sink, prep, round);
|
||||
// Remaining bit triples: produce a single predicate bit via XOR of
|
||||
// consumed AND finishes (placeholder open of zeros for unused slots).
|
||||
const unsigned width = in.aux;
|
||||
std::uint8_t pred = 0;
|
||||
for (unsigned i = 0; i < width; ++i)
|
||||
{
|
||||
if (prep.limb() == 0)
|
||||
break;
|
||||
// Consume one AND round of zeros to keep slot alignment when
|
||||
// triples remain; otherwise skip.
|
||||
try
|
||||
{
|
||||
auto t = prep.take_bit();
|
||||
std::uint8_t mine[2] = {t.a, t.b};
|
||||
sink.submit(round, 0, mine, 2);
|
||||
sink.flush();
|
||||
wait_peer(sink, round);
|
||||
std::uint8_t peer[2]{};
|
||||
sink.read_peer(round, 0, peer, 2);
|
||||
++round;
|
||||
const std::uint8_t d =
|
||||
static_cast<std::uint8_t>(mine[0] ^ peer[0]);
|
||||
const std::uint8_t e =
|
||||
static_cast<std::uint8_t>(mine[1] ^ peer[1]);
|
||||
pred = static_cast<std::uint8_t>(
|
||||
pred ^ edabit::and_finish(t, d, e, id_));
|
||||
}
|
||||
catch (...)
|
||||
{
|
||||
break;
|
||||
}
|
||||
}
|
||||
return pred;
|
||||
}
|
||||
|
||||
std::uint64_t trunc_exact_op(const circuit::inst & in, net::RoundSink & sink,
|
||||
prep::cursor & prep, std::uint16_t & round)
|
||||
{
|
||||
const unsigned n = in.aux;
|
||||
const unsigned s = in.aux2;
|
||||
(void)n;
|
||||
// Open x - r using dabits as r shares.
|
||||
std::uint64_t r = 0;
|
||||
for (unsigned i = 0; i < s; ++i)
|
||||
{
|
||||
auto d = prep.take_dabit();
|
||||
r = (r + (d.arith << i)) & mask_;
|
||||
}
|
||||
const auto delta = exch_sum((words_[in.a] - r) & mask_, sink, round);
|
||||
std::uint8_t wrap = 0;
|
||||
for (unsigned i = 0; i < s; ++i)
|
||||
{
|
||||
auto t = prep.take_bit();
|
||||
std::uint8_t mine[2] = {
|
||||
static_cast<std::uint8_t>(t.a),
|
||||
static_cast<std::uint8_t>(t.b)};
|
||||
sink.submit(round, 0, mine, 2);
|
||||
sink.flush();
|
||||
wait_peer(sink, round);
|
||||
std::uint8_t peer[2]{};
|
||||
sink.read_peer(round, 0, peer, 2);
|
||||
++round;
|
||||
const std::uint8_t d = static_cast<std::uint8_t>(mine[0] ^ peer[0]);
|
||||
const std::uint8_t e = static_cast<std::uint8_t>(mine[1] ^ peer[1]);
|
||||
wrap = static_cast<std::uint8_t>(
|
||||
wrap ^ edabit::and_finish(t, d, e, id_));
|
||||
}
|
||||
return ((delta >> s) + wrap) & mask_;
|
||||
}
|
||||
|
||||
std::uint64_t mux_op(const circuit::inst & in, net::RoundSink & sink,
|
||||
prep::cursor & prep, std::uint16_t & round)
|
||||
{
|
||||
// b + sel*(a-b): Beaver mul of sel and (a-b).
|
||||
const std::uint32_t b_id = in.aux;
|
||||
const auto diff = (words_[in.b] - words_[b_id]) & mask_;
|
||||
std::uint8_t ab[8]{}, bb[8]{}, cb[8]{};
|
||||
prep.take_ring(ab, bb, cb);
|
||||
(void)prep.take_bit(); // sel bit pad reserved in demand
|
||||
std::uint64_t a = 0, b = 0, c = 0;
|
||||
std::memcpy(&a, ab, limb_);
|
||||
std::memcpy(&b, bb, limb_);
|
||||
std::memcpy(&c, cb, limb_);
|
||||
const auto sel = words_[in.a] & 1u;
|
||||
const auto d = exch_sum((sel - a) & mask_, sink, round);
|
||||
const auto e = exch_sum((diff - b) & mask_, sink, round);
|
||||
std::uint64_t z = (c + d * b + e * a) & mask_;
|
||||
if (id_ == 0)
|
||||
z = (z + d * e) & mask_;
|
||||
return (words_[b_id] + z) & mask_;
|
||||
}
|
||||
|
||||
const circuit * circ_ = nullptr;
|
||||
unsigned id_ = 0;
|
||||
std::uint16_t limb_ = 8;
|
||||
std::uint64_t mask_ = ~std::uint64_t{0};
|
||||
std::vector<std::uint64_t> words_;
|
||||
std::vector<std::uint64_t> priv_;
|
||||
std::vector<std::uint8_t> priv_set_;
|
||||
std::chrono::milliseconds wait_{30000};
|
||||
};
|
||||
|
||||
} // namespace mpc
|
||||
} // namespace dpf
|
||||
|
||||
#endif
|
||||
|
|
@ -57,6 +57,13 @@ template <typename T>
|
|||
struct is_vec_tag<T, std::void_t<decltype(T::dpf_vec)>>
|
||||
: std::bool_constant<T::dpf_vec> {};
|
||||
|
||||
template <typename T, typename = void>
|
||||
struct has_from_seed : std::false_type {};
|
||||
template <typename T>
|
||||
struct has_from_seed<T, std::void_t<decltype(
|
||||
T::from_seed(static_cast<const void *>(nullptr), std::size_t{0}))>>
|
||||
: std::true_type {};
|
||||
|
||||
template <typename T, typename = void>
|
||||
struct has_integral_representation : std::false_type {};
|
||||
template <typename T>
|
||||
|
|
@ -542,6 +549,221 @@ group_elem group_final_cw(const typename PRG::block_type & s0,
|
|||
group_add(group_add(group_add(c1, group_neg(c0)), group_neg(Va)), on_path));
|
||||
}
|
||||
|
||||
/// @brief Type-erased group for a comparison payload that supplies `from_seed`,
|
||||
/// `operator+`, and unary `operator-`. The element is at most 256 bytes.
|
||||
struct payload_ops
|
||||
{
|
||||
static constexpr std::size_t cap = 256;
|
||||
std::size_t size = 0;
|
||||
void (*add)(unsigned char *, const unsigned char *, const unsigned char *) = nullptr;
|
||||
void (*neg)(unsigned char *, const unsigned char *) = nullptr;
|
||||
void (*from_node)(unsigned char *, const void *, std::size_t) = nullptr;
|
||||
void (*scale)(unsigned char *, const unsigned char *, std::int64_t) = nullptr;
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
void payload_add(unsigned char * dst, const unsigned char * a, const unsigned char * b)
|
||||
{
|
||||
T x{}, y{};
|
||||
std::memcpy(&x, a, sizeof(T));
|
||||
std::memcpy(&y, b, sizeof(T));
|
||||
const T z = x + y;
|
||||
std::memset(dst, 0, payload_ops::cap);
|
||||
std::memcpy(dst, &z, sizeof(T));
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
void payload_neg(unsigned char * dst, const unsigned char * a)
|
||||
{
|
||||
T x{};
|
||||
std::memcpy(&x, a, sizeof(T));
|
||||
const T z = -x;
|
||||
std::memset(dst, 0, payload_ops::cap);
|
||||
std::memcpy(dst, &z, sizeof(T));
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
void payload_from_node(unsigned char * dst, const void * node, std::size_t n)
|
||||
{
|
||||
const T z = T::from_seed(node, n);
|
||||
std::memset(dst, 0, payload_ops::cap);
|
||||
std::memcpy(dst, &z, sizeof(T));
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
void payload_scale(unsigned char * dst, const unsigned char * a, std::int64_t k)
|
||||
{
|
||||
T g{};
|
||||
std::memcpy(&g, a, sizeof(T));
|
||||
T r{};
|
||||
if (k < 0)
|
||||
{
|
||||
g = -g;
|
||||
k = -k;
|
||||
}
|
||||
auto m = static_cast<std::uint64_t>(k);
|
||||
while (m != 0)
|
||||
{
|
||||
if ((m & 1u) != 0)
|
||||
r = r + g;
|
||||
m >>= 1;
|
||||
if (m != 0)
|
||||
g = g + g;
|
||||
}
|
||||
std::memset(dst, 0, payload_ops::cap);
|
||||
std::memcpy(dst, &r, sizeof(T));
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
payload_ops make_payload_ops()
|
||||
{
|
||||
static_assert(sizeof(T) <= payload_ops::cap,
|
||||
"comparison payload exceeds 256 bytes");
|
||||
static_assert(has_from_seed<T>::value,
|
||||
"comparison payload needs from_seed");
|
||||
payload_ops ops;
|
||||
ops.size = sizeof(T);
|
||||
ops.add = &payload_add<T>;
|
||||
ops.neg = &payload_neg<T>;
|
||||
ops.from_node = &payload_from_node<T>;
|
||||
ops.scale = &payload_scale<T>;
|
||||
return ops;
|
||||
}
|
||||
|
||||
inline void payload_copy(unsigned char * dst, const unsigned char * src, std::size_t n)
|
||||
{
|
||||
std::memset(dst, 0, payload_ops::cap);
|
||||
if (n != 0)
|
||||
std::memcpy(dst, src, n);
|
||||
}
|
||||
|
||||
inline void payload_sgn(const payload_ops & ops, unsigned char * dst,
|
||||
const unsigned char * src, bool neg)
|
||||
{
|
||||
if (!neg)
|
||||
payload_copy(dst, src, ops.size);
|
||||
else
|
||||
ops.neg(dst, src);
|
||||
}
|
||||
|
||||
/// @brief One comparison-level correction in `ops`'s group.
|
||||
/// @details Mirrors `group_value_cw`. When `beta` is null the correction is split
|
||||
/// into a group element independent of δ and an integer coefficient of δ.
|
||||
inline void payload_value_cw(const payload_ops & ops,
|
||||
const void * n0l, const void * n0r, const void * n1l, const void * n1r,
|
||||
std::size_t node_len, std::uint8_t t1, int ai,
|
||||
unsigned char * va, std::int64_t & va_c,
|
||||
const unsigned char * beta, unsigned char * vcw, std::int64_t * coeff)
|
||||
{
|
||||
unsigned char v0l[payload_ops::cap]{}, v0r[payload_ops::cap]{};
|
||||
unsigned char v1l[payload_ops::cap]{}, v1r[payload_ops::cap]{};
|
||||
ops.from_node(v0l, n0l, node_len);
|
||||
ops.from_node(v0r, n0r, node_len);
|
||||
ops.from_node(v1l, n1l, node_len);
|
||||
ops.from_node(v1r, n1r, node_len);
|
||||
const unsigned char * v0k = ai == 0 ? v0l : v0r;
|
||||
const unsigned char * v1k = ai == 0 ? v1l : v1r;
|
||||
const unsigned char * v0lo = ai == 0 ? v0r : v0l;
|
||||
const unsigned char * v1lo = ai == 0 ? v1r : v1l;
|
||||
unsigned char neg_v0[payload_ops::cap]{}, neg_va[payload_ops::cap]{};
|
||||
ops.neg(neg_v0, v0lo);
|
||||
ops.neg(neg_va, va);
|
||||
unsigned char sum[payload_ops::cap]{}, inner[payload_ops::cap]{};
|
||||
ops.add(sum, v1lo, neg_v0);
|
||||
ops.add(inner, sum, neg_va);
|
||||
std::int64_t inner_c = -va_c;
|
||||
if (ai == 1)
|
||||
{
|
||||
if (beta != nullptr)
|
||||
ops.add(inner, inner, beta);
|
||||
else
|
||||
inner_c += 1;
|
||||
}
|
||||
payload_sgn(ops, vcw, inner, t1 != 0);
|
||||
const std::int64_t vcw_c = (t1 != 0) ? -inner_c : inner_c;
|
||||
unsigned char sgn_vcw[payload_ops::cap]{};
|
||||
payload_sgn(ops, sgn_vcw, vcw, t1 != 0);
|
||||
unsigned char neg_v1k[payload_ops::cap]{}, acc[payload_ops::cap]{};
|
||||
ops.neg(neg_v1k, v1k);
|
||||
ops.add(acc, va, neg_v1k);
|
||||
ops.add(acc, acc, v0k);
|
||||
ops.add(va, acc, sgn_vcw);
|
||||
va_c += (t1 != 0) ? -vcw_c : vcw_c;
|
||||
if (coeff != nullptr)
|
||||
*coeff = vcw_c;
|
||||
if (beta != nullptr && vcw_c != 0)
|
||||
{
|
||||
unsigned char extra[payload_ops::cap]{};
|
||||
ops.scale(extra, beta, vcw_c);
|
||||
ops.add(vcw, vcw, extra);
|
||||
}
|
||||
}
|
||||
|
||||
inline void payload_final_cw(const payload_ops & ops,
|
||||
const void * s0, const void * s1, std::size_t node_len, std::uint8_t t1,
|
||||
const unsigned char * va, std::int64_t va_c,
|
||||
const unsigned char * on_path, int on_c,
|
||||
unsigned char * out, std::int64_t * coeff)
|
||||
{
|
||||
unsigned char c0[payload_ops::cap]{}, c1[payload_ops::cap]{};
|
||||
ops.from_node(c0, s0, node_len);
|
||||
ops.from_node(c1, s1, node_len);
|
||||
unsigned char neg_c0[payload_ops::cap]{}, neg_va[payload_ops::cap]{};
|
||||
ops.neg(neg_c0, c0);
|
||||
ops.neg(neg_va, va);
|
||||
unsigned char sum[payload_ops::cap]{}, inner[payload_ops::cap]{};
|
||||
ops.add(sum, c1, neg_c0);
|
||||
ops.add(inner, sum, neg_va);
|
||||
if (on_path != nullptr)
|
||||
ops.add(inner, inner, on_path);
|
||||
std::int64_t inner_c = -va_c + on_c;
|
||||
payload_sgn(ops, out, inner, t1 != 0);
|
||||
const std::int64_t out_c = (t1 != 0) ? -inner_c : inner_c;
|
||||
if (coeff != nullptr)
|
||||
*coeff = out_c;
|
||||
}
|
||||
|
||||
template <typename Word>
|
||||
Word payload_to_word(const unsigned char * bytes, std::size_t n)
|
||||
{
|
||||
Word w{};
|
||||
std::memcpy(&w, bytes, n < sizeof(Word) ? n : sizeof(Word));
|
||||
return w;
|
||||
}
|
||||
|
||||
template <typename T, typename = void>
|
||||
struct payload_has_canonicalize : std::false_type {};
|
||||
|
||||
template <typename T>
|
||||
struct payload_has_canonicalize<T,
|
||||
std::void_t<decltype(T::canonicalize(T{}))>> : std::true_type {};
|
||||
|
||||
template <typename T, typename Word>
|
||||
T payload_from_word(const Word & w)
|
||||
{
|
||||
T t{};
|
||||
utils::raw_memcpy(&t, &w, sizeof(T) < sizeof(Word) ? sizeof(T) : sizeof(Word));
|
||||
if constexpr (payload_has_canonicalize<T>::value)
|
||||
return T::canonicalize(t);
|
||||
return t;
|
||||
}
|
||||
|
||||
template <typename Word>
|
||||
std::int64_t payload_coeff_of(const Word & w)
|
||||
{
|
||||
std::int64_t k = 0;
|
||||
std::memcpy(&k, &w, sizeof(k) < sizeof(Word) ? sizeof(k) : sizeof(Word));
|
||||
return k;
|
||||
}
|
||||
|
||||
template <typename Word>
|
||||
Word payload_coeff_word(std::int64_t k)
|
||||
{
|
||||
Word w{};
|
||||
std::memcpy(&w, &k, sizeof(k) < sizeof(Word) ? sizeof(k) : sizeof(Word));
|
||||
return w;
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
} // namespace dpf
|
||||
|
||||
|
|
|
|||
911
include/dpf/cohort.hpp
Normal file
911
include/dpf/cohort.hpp
Normal file
|
|
@ -0,0 +1,911 @@
|
|||
/// @file dpf/cohort.hpp
|
||||
/// @brief Many classic DPFs that share one public schedule.
|
||||
/// @details Generation is the ordinary special-path walk, run level by level
|
||||
/// across keys so the path bit is read once and `expand_x4` sees
|
||||
/// contiguous seeds. Evaluation uses that same idea on a point, a
|
||||
/// closed interval, or one `sequence_recipe`: correction words are
|
||||
/// hoisted once per level, and interior nodes stay key-interleaved
|
||||
/// (`position * stride + key`) down to the leaves.
|
||||
///
|
||||
/// Leaf layout, with `n` keys:
|
||||
/// - point: `out[key]`
|
||||
/// - interval: leaf node `j`, lane `p` of key `k` is
|
||||
/// `out[(j * n + k) * outputs_per_leaf + p]`
|
||||
/// (same lane order as one-key `eval_interval`)
|
||||
/// - sequence: listed point `q` of key `k` is `out[q * n + k]`
|
||||
///
|
||||
/// `cohort_index(j, k, n)` is `j * n + k`.
|
||||
/// Inner products use the same order and the same weights for every
|
||||
/// key. Generation is the classic (not incremental) dealer key: one
|
||||
/// plaintext point, one payload or `std::tuple` of payloads per key.
|
||||
/// @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_COHORT_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_COHORT_HPP__
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include <array>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <limits>
|
||||
#include <stdexcept>
|
||||
#include <tuple>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/aligned_allocator.hpp"
|
||||
#include "dpf/dpf_key.hpp"
|
||||
#include "dpf/eval_common.hpp"
|
||||
#include "dpf/leaf_node.hpp"
|
||||
#include "dpf/random.hpp"
|
||||
#include "dpf/sequence_recipe.hpp"
|
||||
#include "dpf/utils.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
/// @brief `j * nkeys + key`. Point, interval leaf, and sequence point all use this.
|
||||
HEDLEY_CONST
|
||||
HEDLEY_NO_THROW
|
||||
inline constexpr std::size_t cohort_index(std::size_t j, std::size_t key,
|
||||
std::size_t nkeys) noexcept
|
||||
{
|
||||
return j * nkeys + key;
|
||||
}
|
||||
|
||||
/// @brief Two planes of key-interleaved interior nodes, reused across calls.
|
||||
/// @tparam Key a DPF key or a `party_key` of one. Every key in a walk has this type.
|
||||
template <typename Key>
|
||||
class cohort_scratch
|
||||
{
|
||||
public:
|
||||
using key_type = unwrap_party_key_t<std::decay_t<Key>>;
|
||||
using node_type = typename key_type::interior_node;
|
||||
using allocator = aligned_allocator<node_type>;
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
std::size_t count() const noexcept { return count_; }
|
||||
HEDLEY_NO_THROW
|
||||
std::size_t stride() const noexcept { return stride_; }
|
||||
|
||||
/// @brief Grow so each of `positions` tree slots can hold `nkeys` nodes,
|
||||
/// padded to a multiple of 4 for `expand_x4`.
|
||||
void fit(std::size_t nkeys, std::size_t positions)
|
||||
{
|
||||
count_ = nkeys;
|
||||
stride_ = nkeys == 0 ? 0 : ((nkeys + 3u) & ~std::size_t{3});
|
||||
const std::size_t pos = positions == 0 ? 1 : positions;
|
||||
const std::size_t need = 2 * pos * stride_;
|
||||
if (buf_.size() < need)
|
||||
buf_.resize(need);
|
||||
positions_ = pos;
|
||||
if (cw0_.size() < nkeys)
|
||||
{
|
||||
cw0_.resize(nkeys);
|
||||
cw1_.resize(nkeys);
|
||||
}
|
||||
}
|
||||
|
||||
node_type * plane(int which) noexcept
|
||||
{
|
||||
return buf_.data() + static_cast<std::size_t>(which) * positions_ * stride_;
|
||||
}
|
||||
node_type * cw0() noexcept { return cw0_.data(); }
|
||||
node_type * cw1() noexcept { return cw1_.data(); }
|
||||
|
||||
private:
|
||||
std::size_t count_ = 0;
|
||||
std::size_t stride_ = 0;
|
||||
std::size_t positions_ = 0;
|
||||
std::vector<node_type, allocator> buf_;
|
||||
std::vector<node_type, allocator> cw0_;
|
||||
std::vector<node_type, allocator> cw1_;
|
||||
};
|
||||
|
||||
namespace cohort_detail
|
||||
{
|
||||
|
||||
template <typename T>
|
||||
struct is_std_tuple : std::false_type {};
|
||||
template <typename... Ts>
|
||||
struct is_std_tuple<std::tuple<Ts...>> : std::true_type {};
|
||||
|
||||
template <typename Tree, typename Node>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void expand_keys(const Node * HEDLEY_RESTRICT parents, std::size_t nkeys,
|
||||
bool is_last, const Node * HEDLEY_RESTRICT cw_left,
|
||||
const Node * HEDLEY_RESTRICT cw_right, int which,
|
||||
Node * HEDLEY_RESTRICT dest_left, Node * HEDLEY_RESTRICT dest_right)
|
||||
{
|
||||
// which: 0 left, 1 right, 2 both. Both destinations are real pointers.
|
||||
std::size_t k = 0;
|
||||
for (; k + 4 <= nkeys; k += 4)
|
||||
{
|
||||
alignas(Node) Node left[4];
|
||||
alignas(Node) Node right[4];
|
||||
Tree::expand_x4(parents + k, left, right, is_last);
|
||||
DPF_UNROLL_LOOP
|
||||
for (std::size_t t = 0; t < 4; ++t)
|
||||
{
|
||||
if (which != 1)
|
||||
{
|
||||
dest_left[k + t] = dpf::xor_if_lo_bit(left[t], cw_left[k + t],
|
||||
parents[k + t]);
|
||||
}
|
||||
if (which != 0)
|
||||
{
|
||||
dest_right[k + t] = dpf::xor_if_lo_bit(right[t], cw_right[k + t],
|
||||
parents[k + t]);
|
||||
}
|
||||
}
|
||||
}
|
||||
for (; k < nkeys; ++k)
|
||||
{
|
||||
const auto kids = Tree::expand(parents[k], is_last);
|
||||
if (which != 1)
|
||||
dest_left[k] = dpf::xor_if_lo_bit(kids[0], cw_left[k], parents[k]);
|
||||
if (which != 0)
|
||||
dest_right[k] = dpf::xor_if_lo_bit(kids[1], cw_right[k], parents[k]);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename Tree, typename Node>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void expand_one(const Node * HEDLEY_RESTRICT parents, std::size_t nkeys,
|
||||
bool is_last, const Node * HEDLEY_RESTRICT cw, bool right,
|
||||
Node * HEDLEY_RESTRICT dest)
|
||||
{
|
||||
Node discard;
|
||||
if (right)
|
||||
expand_keys<Tree>(parents, nkeys, is_last, cw, cw, 1, &discard, dest);
|
||||
else
|
||||
expand_keys<Tree>(parents, nkeys, is_last, cw, cw, 0, dest, &discard);
|
||||
}
|
||||
|
||||
template <typename Range>
|
||||
void require_keys(const Range & keys)
|
||||
{
|
||||
if (keys.size() == 0)
|
||||
throw std::invalid_argument("cohort: no keys");
|
||||
}
|
||||
|
||||
template <typename Key, typename Integral>
|
||||
std::size_t nodes_at_level(std::size_t level, Integral from_node, Integral to_node)
|
||||
{
|
||||
const std::size_t offset = Key::depth - level;
|
||||
return static_cast<std::size_t>(
|
||||
utils::shift_right(static_cast<Integral>(to_node - Integral{1}), offset)
|
||||
- utils::shift_right(from_node, offset)) + 1;
|
||||
}
|
||||
|
||||
template <typename NodeT, typename Output, typename Leaf>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
Output lane_value(const Leaf & leaf, std::size_t lane)
|
||||
{
|
||||
if constexpr (utils::is_packed_subbyte_v<Output>)
|
||||
{
|
||||
return extract_leaf<NodeT, Output>(leaf, lane);
|
||||
}
|
||||
else
|
||||
{
|
||||
Output val;
|
||||
std::memcpy(&val,
|
||||
reinterpret_cast<const unsigned char *>(std::addressof(leaf))
|
||||
+ lane * sizeof(Output),
|
||||
sizeof(Output));
|
||||
return val;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename Output, typename W>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
Output mac_add(Output acc, Output val, W && w)
|
||||
{
|
||||
if constexpr (std::is_integral_v<Output> && !std::is_same_v<Output, bool>)
|
||||
{
|
||||
using U = std::make_unsigned_t<Output>;
|
||||
return static_cast<Output>(
|
||||
static_cast<U>(acc)
|
||||
+ static_cast<U>(val) * static_cast<U>(w));
|
||||
}
|
||||
else
|
||||
{
|
||||
return static_cast<Output>(acc + val * Output(w));
|
||||
}
|
||||
}
|
||||
|
||||
template <typename Out>
|
||||
void ensure_size(Out & out, std::size_t n)
|
||||
{
|
||||
if (out.size() < n)
|
||||
out.resize(n);
|
||||
}
|
||||
|
||||
template <std::size_t I, typename Range, typename Input, typename Scratch, typename Fn>
|
||||
void walk_point(const Range & keys, Input walk_x, Input lane_x,
|
||||
Scratch & scratch, Fn && fn)
|
||||
{
|
||||
using key_type = unwrap_party_key_t<std::decay_t<decltype(keys[0])>>;
|
||||
using node = typename key_type::interior_node;
|
||||
using tree = typename key_type::tree;
|
||||
using output = typename key_type::template concrete_output_type<I>;
|
||||
const std::size_t n = keys.size();
|
||||
scratch.fit(n, 1);
|
||||
int src = 0;
|
||||
auto * cur = scratch.plane(src);
|
||||
for (std::size_t k = 0; k < n; ++k)
|
||||
cur[k] = keys[k].root();
|
||||
|
||||
auto mask = key_type::msb_mask;
|
||||
for (std::size_t level = 0; level < key_type::depth; ++level, mask >>= 1)
|
||||
{
|
||||
const bool bit = !!(mask & walk_x);
|
||||
const bool is_last = tree::is_last_level(level, key_type::depth);
|
||||
node * cw = bit ? scratch.cw1() : scratch.cw0();
|
||||
for (std::size_t k = 0; k < n; ++k)
|
||||
cw[k] = keys[k].correction_word(level, bit);
|
||||
const int dst = 1 - src;
|
||||
node * next = scratch.plane(dst);
|
||||
if (bit)
|
||||
{
|
||||
expand_one<tree>(cur, n, is_last, cw, true, next);
|
||||
}
|
||||
else
|
||||
{
|
||||
expand_one<tree>(cur, n, is_last, cw, false, next);
|
||||
}
|
||||
src = dst;
|
||||
cur = next;
|
||||
}
|
||||
|
||||
for (std::size_t k = 0; k < n; ++k)
|
||||
{
|
||||
auto leaf = keys[k].template traverse_exterior<I>(cur[k]);
|
||||
auto handle = make_dpf_output<output>(leaf, lane_x);
|
||||
fn(k, static_cast<output>(handle));
|
||||
}
|
||||
}
|
||||
|
||||
template <std::size_t I, typename Range, typename Integral, typename Scratch, typename Fn>
|
||||
void walk_interval_segment(const Range & keys, Integral from_node, Integral to_node,
|
||||
std::size_t leaf_base, Scratch & scratch, Fn && on_leaf)
|
||||
{
|
||||
using key_type = unwrap_party_key_t<std::decay_t<decltype(keys[0])>>;
|
||||
using node = typename key_type::interior_node;
|
||||
using tree = typename key_type::tree;
|
||||
const std::size_t n = keys.size();
|
||||
constexpr std::size_t depth = key_type::depth;
|
||||
|
||||
std::size_t widest = 1;
|
||||
if constexpr (depth == 0)
|
||||
{
|
||||
widest = nodes_at_level<key_type>(0, from_node, to_node);
|
||||
}
|
||||
else
|
||||
{
|
||||
for (std::size_t level = 1; level <= depth; ++level)
|
||||
{
|
||||
widest = std::max(widest,
|
||||
nodes_at_level<key_type>(level, from_node, to_node));
|
||||
}
|
||||
}
|
||||
scratch.fit(n, widest);
|
||||
const std::size_t stride = scratch.stride();
|
||||
|
||||
int src = 0;
|
||||
for (std::size_t k = 0; k < n; ++k)
|
||||
scratch.plane(0)[k] = keys[k].root();
|
||||
|
||||
for (std::size_t level = 1; level <= depth; ++level)
|
||||
{
|
||||
const std::size_t child_nodes
|
||||
= nodes_at_level<key_type>(level, from_node, to_node);
|
||||
const auto mask = utils::get_node_mask<key_type>(key_type::msb_mask, level);
|
||||
const bool from_offset = static_cast<bool>(mask & from_node);
|
||||
const bool to_offset = from_offset ^ static_cast<bool>(child_nodes & 1u);
|
||||
const bool is_last = tree::is_last_level(level - 1, depth);
|
||||
node * cw_l = scratch.cw0();
|
||||
node * cw_r = scratch.cw1();
|
||||
for (std::size_t k = 0; k < n; ++k)
|
||||
{
|
||||
cw_l[k] = keys[k].correction_word(level - 1, false);
|
||||
cw_r[k] = keys[k].correction_word(level - 1, true);
|
||||
}
|
||||
const int dst = 1 - src;
|
||||
node * prev = scratch.plane(src);
|
||||
node * curr = scratch.plane(dst);
|
||||
std::size_t i = 0;
|
||||
std::size_t j = 0;
|
||||
if (from_offset)
|
||||
{
|
||||
expand_one<tree>(prev + j * stride, n, is_last, cw_r, true,
|
||||
curr + i * stride);
|
||||
++i;
|
||||
++j;
|
||||
}
|
||||
const std::size_t both_end = child_nodes - static_cast<std::size_t>(to_offset);
|
||||
while (i < both_end)
|
||||
{
|
||||
expand_keys<tree>(prev + j * stride, n, is_last, cw_l, cw_r, 2,
|
||||
curr + i * stride, curr + (i + 1) * stride);
|
||||
i += 2;
|
||||
++j;
|
||||
}
|
||||
if (to_offset)
|
||||
{
|
||||
expand_one<tree>(prev + j * stride, n, is_last, cw_l, false,
|
||||
curr + i * stride);
|
||||
}
|
||||
src = dst;
|
||||
}
|
||||
|
||||
const std::size_t leaves = (depth == 0)
|
||||
? nodes_at_level<key_type>(0, from_node, to_node)
|
||||
: nodes_at_level<key_type>(depth, from_node, to_node);
|
||||
node * leaf_plane = scratch.plane(src);
|
||||
for (std::size_t j = 0; j < leaves; ++j)
|
||||
{
|
||||
for (std::size_t k = 0; k < n; ++k)
|
||||
{
|
||||
auto leaf = keys[k].template traverse_exterior<I>(
|
||||
leaf_plane[j * stride + k]);
|
||||
on_leaf(leaf_base + j, k, leaf);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <std::size_t I, typename Range, typename Scratch, typename Fn>
|
||||
void walk_recipe(const Range & keys, const sequence_recipe & recipe,
|
||||
Scratch & scratch, Fn && on_point)
|
||||
{
|
||||
using key_type = unwrap_party_key_t<std::decay_t<decltype(keys[0])>>;
|
||||
using node = typename key_type::interior_node;
|
||||
using tree = typename key_type::tree;
|
||||
using output = typename key_type::template concrete_output_type<I>;
|
||||
if (recipe.depth() != key_type::depth)
|
||||
throw std::invalid_argument("cohort: recipe depth does not match the keys");
|
||||
const std::size_t n = keys.size();
|
||||
const auto nout = recipe.output_indices().size();
|
||||
if (nout == 0 || recipe.num_leaf_nodes() == 0)
|
||||
return;
|
||||
|
||||
scratch.fit(n, std::max(recipe.num_leaf_nodes(), std::size_t{1}));
|
||||
const std::size_t stride = scratch.stride();
|
||||
int src = 0;
|
||||
for (std::size_t k = 0; k < n; ++k)
|
||||
scratch.plane(0)[k] = keys[k].root();
|
||||
|
||||
std::size_t step = 0;
|
||||
for (std::size_t level = 1; level <= key_type::depth; ++level)
|
||||
{
|
||||
const std::size_t step_end = recipe.level_endpoints()[level];
|
||||
const bool is_last = tree::is_last_level(level - 1, key_type::depth);
|
||||
node * cw_l = scratch.cw0();
|
||||
node * cw_r = scratch.cw1();
|
||||
for (std::size_t k = 0; k < n; ++k)
|
||||
{
|
||||
cw_l[k] = keys[k].correction_word(level - 1, false);
|
||||
cw_r[k] = keys[k].correction_word(level - 1, true);
|
||||
}
|
||||
const int dst = 1 - src;
|
||||
node * prev = scratch.plane(src);
|
||||
node * curr = scratch.plane(dst);
|
||||
std::size_t parent_i = 0;
|
||||
std::size_t out_i = 0;
|
||||
for (; step < step_end; ++step, ++parent_i)
|
||||
{
|
||||
const std::int8_t s = recipe.recipe_steps()[step];
|
||||
const bool left = s > std::int8_t{-1};
|
||||
const bool right = s < std::int8_t{1};
|
||||
node * parent = prev + parent_i * stride;
|
||||
if (left && right)
|
||||
{
|
||||
expand_keys<tree>(parent, n, is_last, cw_l, cw_r, 2,
|
||||
curr + out_i * stride, curr + (out_i + 1) * stride);
|
||||
out_i += 2;
|
||||
}
|
||||
else if (left)
|
||||
{
|
||||
expand_one<tree>(parent, n, is_last, cw_l, false,
|
||||
curr + out_i * stride);
|
||||
++out_i;
|
||||
}
|
||||
else
|
||||
{
|
||||
expand_one<tree>(parent, n, is_last, cw_r, true,
|
||||
curr + out_i * stride);
|
||||
++out_i;
|
||||
}
|
||||
}
|
||||
src = dst;
|
||||
}
|
||||
|
||||
node * leaf_plane = scratch.plane(src);
|
||||
constexpr std::size_t opl = key_type::outputs_per_leaf;
|
||||
const auto & idx = recipe.output_indices();
|
||||
std::size_t cached = std::numeric_limits<std::size_t>::max();
|
||||
std::vector<decltype(keys[0].template traverse_exterior<I>(node{}))> leaves(n);
|
||||
for (std::size_t q = 0; q < nout; ++q)
|
||||
{
|
||||
const std::size_t slot = idx[q];
|
||||
const std::size_t node_i = slot / opl;
|
||||
const std::size_t lane = slot % opl;
|
||||
if (node_i != cached)
|
||||
{
|
||||
for (std::size_t k = 0; k < n; ++k)
|
||||
{
|
||||
leaves[k] = keys[k].template traverse_exterior<I>(
|
||||
leaf_plane[node_i * stride + k]);
|
||||
}
|
||||
cached = node_i;
|
||||
}
|
||||
for (std::size_t k = 0; k < n; ++k)
|
||||
on_point(q, k, lane_value<typename key_type::exterior_node, output>(
|
||||
leaves[k], lane));
|
||||
}
|
||||
}
|
||||
|
||||
template <typename InteriorPRG, typename ExteriorPRG, typename Input, typename Payload>
|
||||
struct cohort_dpf_type
|
||||
{
|
||||
using type = utils::dpf_type_t<InteriorPRG, ExteriorPRG, Input, Payload>;
|
||||
};
|
||||
template <typename InteriorPRG, typename ExteriorPRG, typename Input,
|
||||
typename T0, typename... Ts>
|
||||
struct cohort_key_pack
|
||||
{
|
||||
using type = dpf_key<InteriorPRG, ExteriorPRG, Input, T0, Ts...>;
|
||||
};
|
||||
|
||||
template <typename InteriorPRG, typename ExteriorPRG, typename Input, typename... Ts>
|
||||
struct cohort_dpf_type<InteriorPRG, ExteriorPRG, Input, std::tuple<Ts...>>
|
||||
{
|
||||
using type = typename cohort_key_pack<InteriorPRG, ExteriorPRG, Input, Ts...>::type;
|
||||
};
|
||||
|
||||
} // namespace cohort_detail
|
||||
|
||||
/// @brief Evaluate every key at `x`. `out[key]` is that key's raw share of output `I`.
|
||||
/// \complexity O(n m) PRG calls. m is the number of keys and n is `depth`. One batched expand per level. The scratch holds O(m) nodes.
|
||||
template <std::size_t I = 0,
|
||||
typename Range,
|
||||
typename InputT,
|
||||
typename Out>
|
||||
void eval_point_cohort(const Range & keys, InputT x, Out & out,
|
||||
cohort_scratch<std::decay_t<decltype(keys[0])>> * scratch = nullptr)
|
||||
{
|
||||
cohort_detail::require_keys(keys);
|
||||
using stored = std::decay_t<decltype(keys[0])>;
|
||||
using key_type = unwrap_party_key_t<stored>;
|
||||
using output = typename key_type::template concrete_output_type<I>;
|
||||
cohort_scratch<stored> local;
|
||||
auto & ws = scratch != nullptr ? *scratch : local;
|
||||
auto lane_x = keys[0].offset_x(x);
|
||||
auto walk_x = lane_x;
|
||||
utils::flip_msb_if_signed_integral(walk_x);
|
||||
cohort_detail::ensure_size(out, keys.size());
|
||||
cohort_detail::walk_point<I>(keys, walk_x, lane_x, ws,
|
||||
[&](std::size_t k, output v) { out[k] = v; });
|
||||
}
|
||||
|
||||
/// @brief Evaluate every key on the closed interval `[from, to]`.
|
||||
/// @details Leaves are interleaved. See the file comment for the index.
|
||||
/// \complexity Same node count as one `eval_interval`, times m keys. Each level batches the PRG across keys. Scratch holds O(m L) nodes, L the widest level of the interval.
|
||||
template <std::size_t I = 0,
|
||||
typename Range,
|
||||
typename InputT,
|
||||
typename Out>
|
||||
void eval_interval_cohort(const Range & keys, InputT from, InputT to, Out & out,
|
||||
cohort_scratch<std::decay_t<decltype(keys[0])>> * scratch = nullptr)
|
||||
{
|
||||
cohort_detail::require_keys(keys);
|
||||
using stored = std::decay_t<decltype(keys[0])>;
|
||||
using key_type = unwrap_party_key_t<stored>;
|
||||
using output = typename key_type::template concrete_output_type<I>;
|
||||
using integral = typename key_type::integral_type;
|
||||
cohort_scratch<stored> local;
|
||||
auto & ws = scratch != nullptr ? *scratch : local;
|
||||
|
||||
auto from_x = keys[0].offset_x(from);
|
||||
auto to_x = keys[0].offset_x(to);
|
||||
utils::flip_msb_if_signed_integral(from_x);
|
||||
utils::flip_msb_if_signed_integral(to_x);
|
||||
const integral from_node = utils::get_from_node<key_type>(from_x);
|
||||
const integral to_node = utils::get_to_node<key_type>(to_x);
|
||||
constexpr auto to_int = utils::to_integral_type<decltype(from_x)>{};
|
||||
const bool wraps = utils::interval_wraps(
|
||||
static_cast<integral>(to_int(from_x)),
|
||||
static_cast<integral>(to_int(to_x)),
|
||||
utils::bitlength_of_v<decltype(from_x)>);
|
||||
auto segs = utils::split_leaf_nodes(from_node, to_node, key_type::depth, wraps);
|
||||
constexpr std::size_t opl = key_type::outputs_per_leaf;
|
||||
const std::size_t n = keys.size();
|
||||
cohort_detail::ensure_size(out, segs.total * opl * n);
|
||||
|
||||
std::size_t leaf_base = 0;
|
||||
for (std::size_t s = 0; s < segs.n; ++s)
|
||||
{
|
||||
const auto & seg = segs.seg[s];
|
||||
cohort_detail::walk_interval_segment<I>(keys, seg.from_node,
|
||||
seg.to_node, leaf_base, ws,
|
||||
[&](std::size_t leaf_i, std::size_t k, const auto & leaf) {
|
||||
if constexpr (utils::is_packed_subbyte_v<output>)
|
||||
{
|
||||
for (std::size_t p = 0; p < opl; ++p)
|
||||
{
|
||||
out[(leaf_i * n + k) * opl + p]
|
||||
= cohort_detail::lane_value<
|
||||
typename key_type::exterior_node, output>(leaf, p);
|
||||
}
|
||||
}
|
||||
else
|
||||
{
|
||||
std::memcpy(&out[(leaf_i * n + k) * opl],
|
||||
std::addressof(leaf), sizeof(output) * opl);
|
||||
}
|
||||
});
|
||||
leaf_base += seg.count;
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief `eval_interval_cohort` from `min` through `max`.
|
||||
/// \complexity Same as `eval_interval_cohort` on the full domain.
|
||||
template <std::size_t I = 0, typename Range, typename Out>
|
||||
void eval_full_cohort(const Range & keys, Out & out,
|
||||
cohort_scratch<std::decay_t<decltype(keys[0])>> * scratch = nullptr)
|
||||
{
|
||||
cohort_detail::require_keys(keys);
|
||||
using input = typename unwrap_party_key_t<
|
||||
std::decay_t<decltype(keys[0])>>::input_type;
|
||||
eval_interval_cohort<I>(keys, std::numeric_limits<input>::min(),
|
||||
std::numeric_limits<input>::max(), out, scratch);
|
||||
}
|
||||
|
||||
/// @brief Dot every key's interval leaves with `weights`.
|
||||
/// @details `weights[j]` matches one-key `eval_inner_product`: the j-th lane
|
||||
/// of the covering leaves, not a clipped sub-lane. One sum per key.
|
||||
/// \complexity Same interior walk as `eval_interval_cohort`, plus one multiply-add per lane per key. No leaf buffer.
|
||||
template <std::size_t I = 0, typename Range, typename InputT, typename Weights>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_interval_inner_product_cohort(const Range & keys, InputT from, InputT to,
|
||||
Weights && weights,
|
||||
cohort_scratch<std::decay_t<decltype(keys[0])>> * scratch = nullptr)
|
||||
{
|
||||
cohort_detail::require_keys(keys);
|
||||
using stored = std::decay_t<decltype(keys[0])>;
|
||||
using key_type = unwrap_party_key_t<stored>;
|
||||
using output = typename key_type::template concrete_output_type<I>;
|
||||
using integral = typename key_type::integral_type;
|
||||
cohort_scratch<stored> local;
|
||||
auto & ws = scratch != nullptr ? *scratch : local;
|
||||
const std::size_t n = keys.size();
|
||||
std::vector<output> acc(n);
|
||||
|
||||
auto from_x = keys[0].offset_x(from);
|
||||
auto to_x = keys[0].offset_x(to);
|
||||
utils::flip_msb_if_signed_integral(from_x);
|
||||
utils::flip_msb_if_signed_integral(to_x);
|
||||
const integral from_node = utils::get_from_node<key_type>(from_x);
|
||||
const integral to_node = utils::get_to_node<key_type>(to_x);
|
||||
constexpr auto to_int = utils::to_integral_type<decltype(from_x)>{};
|
||||
const bool wraps = utils::interval_wraps(
|
||||
static_cast<integral>(to_int(from_x)),
|
||||
static_cast<integral>(to_int(to_x)),
|
||||
utils::bitlength_of_v<decltype(from_x)>);
|
||||
auto segs = utils::split_leaf_nodes(from_node, to_node, key_type::depth, wraps);
|
||||
constexpr std::size_t opl = key_type::outputs_per_leaf;
|
||||
|
||||
std::size_t leaf_base = 0;
|
||||
for (std::size_t s = 0; s < segs.n; ++s)
|
||||
{
|
||||
const auto & seg = segs.seg[s];
|
||||
cohort_detail::walk_interval_segment<I>(keys, seg.from_node,
|
||||
seg.to_node, leaf_base, ws,
|
||||
[&](std::size_t leaf_i, std::size_t k, const auto & leaf) {
|
||||
for (std::size_t p = 0; p < opl; ++p)
|
||||
{
|
||||
const auto val = cohort_detail::lane_value<
|
||||
typename key_type::exterior_node, output>(leaf, p);
|
||||
const std::size_t w = leaf_i * opl + p;
|
||||
acc[k] = cohort_detail::mac_add(acc[k], val, weights[w]);
|
||||
}
|
||||
});
|
||||
leaf_base += seg.count;
|
||||
}
|
||||
return acc;
|
||||
}
|
||||
|
||||
/// @brief Evaluate every key on one compiled recipe.
|
||||
/// @details `out[q * n + k]` is key `k` at listed point `q`.
|
||||
/// \complexity One recipe traversal. Each visited node is expanded for all m keys together. Scratch holds O(m L) nodes, L the recipe's leaf count.
|
||||
template <std::size_t I = 0, typename Range, typename Out>
|
||||
void eval_sequence_cohort(const Range & keys, const sequence_recipe & recipe,
|
||||
Out & out, cohort_scratch<std::decay_t<decltype(keys[0])>> * scratch = nullptr)
|
||||
{
|
||||
cohort_detail::require_keys(keys);
|
||||
using stored = std::decay_t<decltype(keys[0])>;
|
||||
using key_type = unwrap_party_key_t<stored>;
|
||||
using output = typename key_type::template concrete_output_type<I>;
|
||||
cohort_scratch<stored> local;
|
||||
auto & ws = scratch != nullptr ? *scratch : local;
|
||||
const std::size_t n = keys.size();
|
||||
cohort_detail::ensure_size(out, recipe.output_indices().size() * n);
|
||||
cohort_detail::walk_recipe<I>(keys, recipe, ws,
|
||||
[&](std::size_t q, std::size_t k, output v) {
|
||||
out[cohort_index(q, k, n)] = v;
|
||||
});
|
||||
}
|
||||
|
||||
/// @brief Compile `[begin, end)` once, then `eval_sequence_cohort` on that recipe.
|
||||
/// @throws std::runtime_error if the range is not sorted nondecreasing.
|
||||
/// \complexity One recipe build, O(k log n) in the point list, then the recipe walk.
|
||||
template <std::size_t I = 0, typename Range, typename ForwardIterator, typename Out>
|
||||
void eval_sequence_cohort(const Range & keys, ForwardIterator begin,
|
||||
ForwardIterator end, Out & out,
|
||||
cohort_scratch<std::decay_t<decltype(keys[0])>> * scratch = nullptr)
|
||||
{
|
||||
cohort_detail::require_keys(keys);
|
||||
using key_type = unwrap_party_key_t<std::decay_t<decltype(keys[0])>>;
|
||||
auto recipe = make_sequence_recipe<key_type>(begin, end);
|
||||
eval_sequence_cohort<I>(keys, recipe, out, scratch);
|
||||
}
|
||||
|
||||
/// @brief Inner product of a recipe's listed points. `weights[q]` pairs with point `q`.
|
||||
/// \complexity Same walk as `eval_sequence_cohort`, plus one multiply-add per point per key.
|
||||
template <std::size_t I = 0, typename Range, typename Weights>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_sequence_inner_product_cohort(const Range & keys,
|
||||
const sequence_recipe & recipe, Weights && weights,
|
||||
cohort_scratch<std::decay_t<decltype(keys[0])>> * scratch = nullptr)
|
||||
{
|
||||
cohort_detail::require_keys(keys);
|
||||
using stored = std::decay_t<decltype(keys[0])>;
|
||||
using key_type = unwrap_party_key_t<stored>;
|
||||
using output = typename key_type::template concrete_output_type<I>;
|
||||
cohort_scratch<stored> local;
|
||||
auto & ws = scratch != nullptr ? *scratch : local;
|
||||
std::vector<output> acc(keys.size());
|
||||
cohort_detail::walk_recipe<I>(keys, recipe, ws,
|
||||
[&](std::size_t q, std::size_t k, output v) {
|
||||
acc[k] = cohort_detail::mac_add(acc[k], v, weights[q]);
|
||||
});
|
||||
return acc;
|
||||
}
|
||||
|
||||
/// @brief Compile `[begin, end)` once, then the recipe inner product.
|
||||
/// @throws std::runtime_error if the range is not sorted nondecreasing.
|
||||
/// \complexity One recipe build plus the recipe inner product.
|
||||
template <std::size_t I = 0, typename Range, typename ForwardIterator, typename Weights>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_sequence_inner_product_cohort(const Range & keys, ForwardIterator begin,
|
||||
ForwardIterator end, Weights && weights,
|
||||
cohort_scratch<std::decay_t<decltype(keys[0])>> * scratch = nullptr)
|
||||
{
|
||||
cohort_detail::require_keys(keys);
|
||||
using key_type = unwrap_party_key_t<std::decay_t<decltype(keys[0])>>;
|
||||
auto recipe = make_sequence_recipe<key_type>(begin, end);
|
||||
return eval_sequence_inner_product_cohort<I>(keys, recipe,
|
||||
std::forward<Weights>(weights), scratch);
|
||||
}
|
||||
|
||||
/// @brief One party's keys plus the scratch those walks reuse.
|
||||
template <typename Key>
|
||||
class cohort
|
||||
{
|
||||
public:
|
||||
using key_type = std::decay_t<Key>;
|
||||
|
||||
cohort() = default;
|
||||
explicit cohort(std::vector<key_type> keys) : keys_(std::move(keys)) {}
|
||||
|
||||
std::size_t size() const noexcept { return keys_.size(); }
|
||||
const std::vector<key_type> & keys() const noexcept { return keys_; }
|
||||
std::vector<key_type> & keys() noexcept { return keys_; }
|
||||
cohort_scratch<key_type> & scratch() noexcept { return scratch_; }
|
||||
|
||||
template <std::size_t I = 0, typename InputT, typename Out>
|
||||
void eval_point(InputT x, Out & out)
|
||||
{
|
||||
eval_point_cohort<I>(keys_, x, out, &scratch_);
|
||||
}
|
||||
|
||||
template <std::size_t I = 0, typename InputT, typename Out>
|
||||
void eval_interval(InputT from, InputT to, Out & out)
|
||||
{
|
||||
eval_interval_cohort<I>(keys_, from, to, out, &scratch_);
|
||||
}
|
||||
|
||||
template <std::size_t I = 0, typename Out>
|
||||
void eval_full(Out & out)
|
||||
{
|
||||
eval_full_cohort<I>(keys_, out, &scratch_);
|
||||
}
|
||||
|
||||
template <std::size_t I = 0, typename InputT, typename Weights>
|
||||
auto eval_interval_inner_product(InputT from, InputT to, Weights && weights)
|
||||
{
|
||||
return eval_interval_inner_product_cohort<I>(keys_, from, to,
|
||||
std::forward<Weights>(weights), &scratch_);
|
||||
}
|
||||
|
||||
template <std::size_t I = 0, typename Out>
|
||||
void eval_sequence(const sequence_recipe & recipe, Out & out)
|
||||
{
|
||||
eval_sequence_cohort<I>(keys_, recipe, out, &scratch_);
|
||||
}
|
||||
|
||||
template <std::size_t I = 0, typename ForwardIterator, typename Out>
|
||||
void eval_sequence(ForwardIterator begin, ForwardIterator end, Out & out)
|
||||
{
|
||||
eval_sequence_cohort<I>(keys_, begin, end, out, &scratch_);
|
||||
}
|
||||
|
||||
template <std::size_t I = 0, typename Weights>
|
||||
auto eval_sequence_inner_product(const sequence_recipe & recipe,
|
||||
Weights && weights)
|
||||
{
|
||||
return eval_sequence_inner_product_cohort<I>(keys_, recipe,
|
||||
std::forward<Weights>(weights), &scratch_);
|
||||
}
|
||||
|
||||
template <std::size_t I = 0, typename ForwardIterator, typename Weights>
|
||||
auto eval_sequence_inner_product(ForwardIterator begin, ForwardIterator end,
|
||||
Weights && weights)
|
||||
{
|
||||
return eval_sequence_inner_product_cohort<I>(keys_, begin, end,
|
||||
std::forward<Weights>(weights), &scratch_);
|
||||
}
|
||||
|
||||
private:
|
||||
std::vector<key_type> keys_;
|
||||
cohort_scratch<key_type> scratch_;
|
||||
};
|
||||
|
||||
/// @brief Classic keys for one point and many payloads.
|
||||
/// @details The path bit is shared. Each level expands every key's seeds with
|
||||
/// `expand_x4`, then writes that key's correction word. Payloads may
|
||||
/// be a single output or a `std::tuple` of outputs. Comparison,
|
||||
/// incremental, verifiable, and extractable tags are not part of this
|
||||
/// walk; build those with `make_dpf` one key at a time.
|
||||
/// @param x plaintext domain point
|
||||
/// @param begin first payload
|
||||
/// @param end past the last payload
|
||||
/// @return party-0 cohort and party-1 cohort, in payload order
|
||||
/// \complexity O(n m) PRG calls. m is the number of payloads and n is `depth`. Roots are sampled first; each level then expands contiguous seeds.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
typename Iter>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto make_dpf_cohort(InputT x, Iter begin, Iter end)
|
||||
{
|
||||
using input_type = std::decay_t<InputT>;
|
||||
static_assert(!is_secret_share_v<input_type>,
|
||||
"make_dpf_cohort: domain point must be plaintext");
|
||||
using payload = std::decay_t<decltype(*begin)>;
|
||||
static_assert(cohort_detail::is_std_tuple<payload>::value
|
||||
|| !is_secret_share_v<payload>,
|
||||
"make_dpf_cohort: payloads must be plaintext");
|
||||
|
||||
std::vector<payload> ys(begin, end);
|
||||
if (ys.empty())
|
||||
throw std::invalid_argument("make_dpf_cohort: no payloads");
|
||||
|
||||
using dpf_type = typename cohort_detail::cohort_dpf_type<
|
||||
InteriorPRG, ExteriorPRG, input_type, payload>::type;
|
||||
using node = typename dpf_type::interior_node;
|
||||
using tree = typename dpf_type::tree;
|
||||
using words = typename dpf_type::correction_words_array;
|
||||
using advice = typename dpf_type::correction_advice_array;
|
||||
using alloc = aligned_allocator<node>;
|
||||
|
||||
utils::flip_msb_if_signed_integral(x);
|
||||
const std::size_t m = ys.size();
|
||||
constexpr std::size_t depth = dpf_type::depth;
|
||||
|
||||
std::vector<node, alloc> s0(m), s1(m), root0(m), root1(m);
|
||||
std::vector<words> cws(m);
|
||||
std::vector<advice> adv(m);
|
||||
|
||||
for (std::size_t k = 0; k < m; ++k)
|
||||
{
|
||||
node tmp[2];
|
||||
tree::root_init(tmp, []() -> node {
|
||||
return static_cast<node>(dpf::uniform_sample<node>());
|
||||
});
|
||||
root0[k] = s0[k] = tmp[0];
|
||||
root1[k] = s1[k] = tmp[1];
|
||||
}
|
||||
|
||||
auto mask = dpf_type::msb_mask;
|
||||
for (std::size_t level = 0; level < depth; ++level, mask >>= 1)
|
||||
{
|
||||
const bool bit = !!(mask & x);
|
||||
const bool is_last = tree::is_last_level(level, depth);
|
||||
std::size_t k = 0;
|
||||
auto step = [&](std::size_t i) {
|
||||
const auto kids0 = tree::expand(s0[i], is_last);
|
||||
const auto kids1 = tree::expand(s1[i], is_last);
|
||||
const bool c0 = static_cast<bool>(dpf::get_lo_bit(s0[i]));
|
||||
const bool c1 = static_cast<bool>(dpf::get_lo_bit(s1[i]));
|
||||
tree::make_cw(cws[i][level], adv[i][level], kids0, kids1,
|
||||
s0[i], s1[i], bit, is_last);
|
||||
const node n0 = tree::advance(s0[i], kids0, cws[i][level],
|
||||
adv[i][level], bit, c0, is_last);
|
||||
const node n1 = tree::advance(s1[i], kids1, cws[i][level],
|
||||
adv[i][level], bit, c1, is_last);
|
||||
s0[i] = n0;
|
||||
s1[i] = n1;
|
||||
};
|
||||
for (; k + 4 <= m; k += 4)
|
||||
{
|
||||
alignas(node) node l0[4], r0[4], l1[4], r1[4];
|
||||
tree::expand_x4(s0.data() + k, l0, r0, is_last);
|
||||
tree::expand_x4(s1.data() + k, l1, r1, is_last);
|
||||
DPF_UNROLL_LOOP
|
||||
for (std::size_t t = 0; t < 4; ++t)
|
||||
{
|
||||
const std::size_t i = k + t;
|
||||
const std::array<node, 2> kids0{l0[t], r0[t]};
|
||||
const std::array<node, 2> kids1{l1[t], r1[t]};
|
||||
const bool c0 = static_cast<bool>(dpf::get_lo_bit(s0[i]));
|
||||
const bool c1 = static_cast<bool>(dpf::get_lo_bit(s1[i]));
|
||||
tree::make_cw(cws[i][level], adv[i][level], kids0, kids1,
|
||||
s0[i], s1[i], bit, is_last);
|
||||
const node n0 = tree::advance(s0[i], kids0, cws[i][level],
|
||||
adv[i][level], bit, c0, is_last);
|
||||
const node n1 = tree::advance(s1[i], kids1, cws[i][level],
|
||||
adv[i][level], bit, c1, is_last);
|
||||
s0[i] = n0;
|
||||
s1[i] = n1;
|
||||
}
|
||||
}
|
||||
for (; k < m; ++k)
|
||||
step(k);
|
||||
}
|
||||
|
||||
using party0 = party_key<0, dpf_type>;
|
||||
using party1 = party_key<1, dpf_type>;
|
||||
std::vector<party0> k0;
|
||||
std::vector<party1> k1;
|
||||
k0.reserve(m);
|
||||
k1.reserve(m);
|
||||
for (std::size_t i = 0; i < m; ++i)
|
||||
{
|
||||
const bool sign0 = static_cast<bool>(dpf::get_lo_bit(s0[i]));
|
||||
const node seed0 = dpf::unset_lo_2bits(s0[i]);
|
||||
const node seed1 = dpf::unset_lo_2bits(s1[i]);
|
||||
auto built = [&]() {
|
||||
if constexpr (cohort_detail::is_std_tuple<payload>::value)
|
||||
{
|
||||
return std::apply([&](const auto & ...p) {
|
||||
return dpf::make_leaves<ExteriorPRG>(x, seed0, seed1, sign0,
|
||||
std::size_t{0}, p...);
|
||||
}, ys[i]);
|
||||
}
|
||||
else
|
||||
{
|
||||
return dpf::make_leaves<ExteriorPRG>(x, seed0, seed1, sign0,
|
||||
std::size_t{0}, ys[i]);
|
||||
}
|
||||
}();
|
||||
auto paired = dpf::make_party_key_pair(
|
||||
dpf_type{root0[i], cws[i], adv[i], built.first.first,
|
||||
built.first.second, input_type{}},
|
||||
dpf_type{root1[i], cws[i], adv[i], built.second.first,
|
||||
built.second.second, input_type{}});
|
||||
k0.push_back(std::move(paired.first));
|
||||
k1.push_back(std::move(paired.second));
|
||||
}
|
||||
return std::make_pair(cohort<party0>(std::move(k0)),
|
||||
cohort<party1>(std::move(k1)));
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_COHORT_HPP__
|
||||
4293
include/dpf/compose.hpp
Normal file
4293
include/dpf/compose.hpp
Normal file
File diff suppressed because it is too large
Load diff
134
include/dpf/compose_async.hpp
Normal file
134
include/dpf/compose_async.hpp
Normal file
|
|
@ -0,0 +1,134 @@
|
|||
/// @file dpf/compose_async.hpp
|
||||
/// @brief Drive compose `plan`s on `async_stream_array`.
|
||||
/// @details The same `drive_via_schedule` path used on sync `stream_array`
|
||||
/// runs on `async_round_sink`, so peer reads are event-driven.
|
||||
/// `drive_options::n_lanes` sizes the stream pool and
|
||||
/// `drive_options::framing` picks the wire layout.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_COMPOSE_ASYNC_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_COMPOSE_ASYNC_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <exception>
|
||||
#include <map>
|
||||
#include <mutex>
|
||||
#include <thread>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/net/asio_ns.hpp"
|
||||
|
||||
#include "dpf/compose.hpp"
|
||||
#include "dpf/net/async_round_sink.hpp"
|
||||
#include "dpf/net/async_stream_array.hpp"
|
||||
#include "dpf/net/round_lane.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace protocol
|
||||
{
|
||||
|
||||
/// @brief Sink options that follow a drive's framing and drain budget.
|
||||
inline net::sink_options sink_options_for(const drive_options & opt)
|
||||
{
|
||||
net::sink_options so;
|
||||
so.framing = opt.framing;
|
||||
if (opt.wait_timeout.count() != 0)
|
||||
so.drain_timeout = opt.wait_timeout;
|
||||
return so;
|
||||
}
|
||||
|
||||
/// @brief Drive a peer-only plan on an `async_stream_array` (overlapped I/O).
|
||||
/// @param instances RoundSink batch width (protocol instances per round).
|
||||
inline void drive_plan_on_async_streams(const plan & p,
|
||||
net::async_stream_array & streams,
|
||||
std::vector<std::vector<std::uint8_t>> & values,
|
||||
const std::map<std::uint32_t, kernel_fn> & kernels, std::size_t party = 0,
|
||||
std::size_t instances = 1, const drive_options & opt = {},
|
||||
net::sink_options sopt = {})
|
||||
{
|
||||
auto slots = p.slot_bytes_all();
|
||||
if (!slots.empty() && streams.size() == 0)
|
||||
throw std::invalid_argument(
|
||||
"drive_plan_on_async_streams: empty stream array");
|
||||
sopt.framing = opt.framing;
|
||||
net::async_round_sink sink(streams, std::move(slots), instances,
|
||||
std::move(sopt));
|
||||
drive_options local = opt;
|
||||
if (local.workers != nullptr && local.pump == nullptr)
|
||||
local.pump = &streams.context();
|
||||
drive_via_schedule(p, sink, values, kernels, party, local);
|
||||
}
|
||||
|
||||
/// @brief Two parties, two threads, split-io async memory pair; drive both plans.
|
||||
/// @details A party that fails closes its end, so the other fails fast.
|
||||
inline void drive_both_on_async_streams(const plan & p0, const plan & p1,
|
||||
std::vector<std::vector<std::uint8_t>> & v0,
|
||||
std::vector<std::vector<std::uint8_t>> & v1,
|
||||
const std::map<std::uint32_t, kernel_fn> & kernels = {},
|
||||
std::size_t instances = 1, const drive_options & opt = {})
|
||||
{
|
||||
const auto slots = p0.slot_bytes_all();
|
||||
if (slots != p1.slot_bytes_all())
|
||||
throw std::invalid_argument(
|
||||
"drive_both_on_async_streams: party slot shapes differ");
|
||||
const std::size_t nstreams =
|
||||
slots.empty() ? 1 : lanes_for_plan(slots.size(), opt);
|
||||
asio::io_context io0;
|
||||
asio::io_context io1;
|
||||
auto peer = net::make_async_dual_memory_stream_pair(io0, io1, nstreams);
|
||||
auto work0 = asio::make_work_guard(io0);
|
||||
auto work1 = asio::make_work_guard(io1);
|
||||
std::mutex err_mu;
|
||||
std::exception_ptr err;
|
||||
auto note = [&](std::exception_ptr e) {
|
||||
std::lock_guard<std::mutex> lock(err_mu);
|
||||
if (!err)
|
||||
err = std::move(e);
|
||||
};
|
||||
std::thread t0([&] {
|
||||
try
|
||||
{
|
||||
drive_plan_on_async_streams(p0, peer.first, v0, kernels, 0,
|
||||
instances, opt);
|
||||
}
|
||||
catch (...)
|
||||
{
|
||||
note(std::current_exception());
|
||||
peer.first.close();
|
||||
}
|
||||
work0.reset();
|
||||
});
|
||||
std::thread t1([&] {
|
||||
try
|
||||
{
|
||||
drive_plan_on_async_streams(p1, peer.second, v1, kernels, 1,
|
||||
instances, opt);
|
||||
}
|
||||
catch (...)
|
||||
{
|
||||
note(std::current_exception());
|
||||
peer.second.close();
|
||||
}
|
||||
work1.reset();
|
||||
});
|
||||
t0.join();
|
||||
t1.join();
|
||||
if (err)
|
||||
std::rethrow_exception(err);
|
||||
}
|
||||
|
||||
/// @brief Same plan on both parties (symmetric compose graphs).
|
||||
inline void drive_both_on_async_streams(const plan & p,
|
||||
std::vector<std::vector<std::uint8_t>> & v0,
|
||||
std::vector<std::vector<std::uint8_t>> & v1,
|
||||
const std::map<std::uint32_t, kernel_fn> & kernels = {},
|
||||
std::size_t instances = 1, const drive_options & opt = {})
|
||||
{
|
||||
drive_both_on_async_streams(p, p, v0, v1, kernels, instances, opt);
|
||||
}
|
||||
|
||||
} // namespace protocol
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_COMPOSE_ASYNC_HPP__
|
||||
|
|
@ -12,6 +12,7 @@
|
|||
#define LIBDPF_INCLUDE_DPF_CONSTRAINED_CMP_HPP__
|
||||
|
||||
#include <cstdint>
|
||||
#include <stdexcept>
|
||||
#include <type_traits>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
|
@ -61,15 +62,18 @@ constexpr void ccmp_party_terms(std::uint64_t x, uint8_t party,
|
|||
}
|
||||
|
||||
/// @brief Opened result of Π_CCMP when both inputs are known (local joint sim).
|
||||
/// @details Precondition: `|x0 - x1| = 1`.
|
||||
/// @details Aborts unless `|x0 - x1| = 1`.
|
||||
/// @param x0 the `x0`
|
||||
/// @param x1 the `x1`
|
||||
/// @return Opened result of Π_CCMP when both inputs are known (local joint sim)
|
||||
HEDLEY_NO_THROW
|
||||
/// @throws std::invalid_argument when the inputs do not differ by one
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_CONST
|
||||
constexpr uint8_t local_ccmp(std::uint64_t x0, std::uint64_t x1) noexcept
|
||||
uint8_t local_ccmp(std::uint64_t x0, std::uint64_t x1)
|
||||
{
|
||||
const std::uint64_t diff = x0 > x1 ? x0 - x1 : x1 - x0;
|
||||
if (diff != 1ULL)
|
||||
throw std::invalid_argument(
|
||||
"constrained comparison: inputs must differ by exactly one");
|
||||
uint8_t z00 = 0, z01 = 0, l0 = 0;
|
||||
uint8_t z10 = 0, z11 = 0, l1 = 0;
|
||||
ccmp_party_terms(x0, 0, z00, z01, l0);
|
||||
|
|
@ -87,11 +91,10 @@ constexpr uint8_t local_ccmp(std::uint64_t x0, std::uint64_t x1) noexcept
|
|||
/// @param x0 the first integer
|
||||
/// @param x1 the second integer
|
||||
/// @return Same as `local_ccmp` for any unsigned or enum-convertible integer
|
||||
/// @throws std::invalid_argument when the inputs do not differ by one
|
||||
template <typename T0, typename T1>
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_CONST
|
||||
constexpr uint8_t local_ccmp_int(T0 x0, T1 x1) noexcept
|
||||
uint8_t local_ccmp_int(T0 x0, T1 x1)
|
||||
{
|
||||
static_assert(std::is_integral_v<T0> && std::is_integral_v<T1>,
|
||||
"local_ccmp_int: integral operands");
|
||||
|
|
|
|||
155
include/dpf/cost_pass.hpp
Normal file
155
include/dpf/cost_pass.hpp
Normal file
|
|
@ -0,0 +1,155 @@
|
|||
/// @file dpf/cost_pass.hpp
|
||||
/// @brief Annotate a recorder/composer plan with conversion strategy choices.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_COST_PASS_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_COST_PASS_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <string>
|
||||
#include <vector>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/beaver.hpp"
|
||||
#include "dpf/compose.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace cost
|
||||
{
|
||||
|
||||
enum class strategy : unsigned char
|
||||
{
|
||||
dcf_mask = 0,
|
||||
edabit_msb = 1,
|
||||
full_a2b_adder = 2,
|
||||
trunc_prob = 3,
|
||||
trunc_exact = 4
|
||||
};
|
||||
|
||||
struct choice
|
||||
{
|
||||
std::uint32_t node_id = 0;
|
||||
strategy pick = strategy::edabit_msb;
|
||||
std::size_t estimated_bytes = 0;
|
||||
std::size_t estimated_rounds = 0;
|
||||
std::string reason;
|
||||
};
|
||||
|
||||
struct report
|
||||
{
|
||||
std::vector<choice> choices;
|
||||
std::size_t setup_bytes = 0;
|
||||
std::size_t online_bytes = 0;
|
||||
beavers::schedule_objective objective = beavers::schedule_objective::rounds;
|
||||
};
|
||||
|
||||
/// @brief Cost model constants (bytes / rounds) for strategy selection.
|
||||
struct model
|
||||
{
|
||||
std::size_t dcf_bytes_per_bit = 16;
|
||||
std::size_t edabit_bytes_per_bit = 8;
|
||||
std::size_t a2b_bytes_per_bit = 24;
|
||||
std::size_t trunc_exact_bytes = 16;
|
||||
std::size_t trunc_prob_bytes = 0;
|
||||
};
|
||||
|
||||
inline strategy pick_compare(beavers::schedule_objective obj, unsigned width,
|
||||
const model & m, choice & out)
|
||||
{
|
||||
const std::size_t dcf = m.dcf_bytes_per_bit * width;
|
||||
const std::size_t eda = m.edabit_bytes_per_bit * width;
|
||||
const std::size_t a2b = m.a2b_bytes_per_bit * width;
|
||||
if (obj == beavers::schedule_objective::rounds)
|
||||
{
|
||||
// Prefer few rounds: DCF-mask (1 open) over full A2B adder.
|
||||
out.pick = strategy::dcf_mask;
|
||||
out.estimated_bytes = dcf;
|
||||
out.estimated_rounds = 1;
|
||||
out.reason = "rounds: dcf_mask";
|
||||
(void)eda;
|
||||
(void)a2b;
|
||||
return out.pick;
|
||||
}
|
||||
// Prep: pick cheapest bytes.
|
||||
if (eda <= dcf && eda <= a2b)
|
||||
{
|
||||
out.pick = strategy::edabit_msb;
|
||||
out.estimated_bytes = eda;
|
||||
out.estimated_rounds = 2;
|
||||
out.reason = "prep: edabit_msb";
|
||||
}
|
||||
else if (dcf <= a2b)
|
||||
{
|
||||
out.pick = strategy::dcf_mask;
|
||||
out.estimated_bytes = dcf;
|
||||
out.estimated_rounds = 1;
|
||||
out.reason = "prep: dcf_mask";
|
||||
}
|
||||
else
|
||||
{
|
||||
out.pick = strategy::full_a2b_adder;
|
||||
out.estimated_bytes = a2b;
|
||||
out.estimated_rounds = static_cast<std::size_t>(width);
|
||||
out.reason = "prep: full_a2b_adder";
|
||||
}
|
||||
return out.pick;
|
||||
}
|
||||
|
||||
inline strategy pick_trunc(beavers::schedule_objective obj, bool need_exact,
|
||||
const model & m, choice & out)
|
||||
{
|
||||
if (!need_exact)
|
||||
{
|
||||
out.pick = strategy::trunc_prob;
|
||||
out.estimated_bytes = m.trunc_prob_bytes;
|
||||
out.estimated_rounds = 0;
|
||||
out.reason = "trunc_prob";
|
||||
return out.pick;
|
||||
}
|
||||
out.pick = strategy::trunc_exact;
|
||||
out.estimated_bytes = m.trunc_exact_bytes;
|
||||
out.estimated_rounds = (obj == beavers::schedule_objective::rounds) ? 1 : 2;
|
||||
out.reason = "trunc_exact";
|
||||
return out.pick;
|
||||
}
|
||||
|
||||
/// @brief Annotate a plan: for each share_cmp / trunc opcode, record a choice.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline report annotate(const protocol::plan & p,
|
||||
beavers::schedule_objective obj = beavers::schedule_objective::rounds,
|
||||
unsigned default_width = 64, bool exact_trunc = true,
|
||||
model m = {})
|
||||
{
|
||||
report r;
|
||||
r.objective = obj;
|
||||
r.setup_bytes = p.setup_bytes();
|
||||
r.online_bytes = p.online_bytes();
|
||||
for (auto n : p.nodes())
|
||||
{
|
||||
const auto op = p.opcode_of(n.id);
|
||||
choice c;
|
||||
c.node_id = n.id;
|
||||
if (op == protocol::opcodes::share_cmp
|
||||
|| op == protocol::opcodes::user_base + 1)
|
||||
{
|
||||
pick_compare(obj, default_width, m, c);
|
||||
r.choices.push_back(c);
|
||||
}
|
||||
else if (op == protocol::opcodes::trunc_exact
|
||||
|| op == protocol::opcodes::trunc_prob
|
||||
|| op == protocol::opcodes::mul_trunc)
|
||||
{
|
||||
pick_trunc(obj, exact_trunc || op != protocol::opcodes::trunc_prob,
|
||||
m, c);
|
||||
r.choices.push_back(c);
|
||||
}
|
||||
}
|
||||
// Empty plans (no compare/trunc ops) get no synthetic choices.
|
||||
return r;
|
||||
}
|
||||
|
||||
} // namespace cost
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_COST_PASS_HPP__
|
||||
|
|
@ -3,6 +3,7 @@
|
|||
/// @details `lt`/`leq`/`gt`/`geq` (+ `_at`) take `(if_true, if_false=0)`.
|
||||
/// Eval walks the same GGM tree as the DPF (per-level value CWs).
|
||||
/// `eq` / `eq_at` are synonyms for ordinary point placements.
|
||||
/// @note Following Boyle, Chandran, Gilboa, Gupta, Ishai, Kumar, and Rathee, EUROCRYPT 2021 (ePrint 2020/1392): a comparison is a DPF spine plus one value-correction word per level. `eq` is a point payload, not that DCF.
|
||||
/// @copyright Copyright (c) 2019-2026 Ryan Henry and [others](@ref authors)
|
||||
/// @license Released under a GNU General Public v2.0 (GPLv2) license.
|
||||
|
||||
|
|
@ -163,7 +164,7 @@ Beta sub_beta(const Beta & a, const Beta & b) noexcept
|
|||
if constexpr (std::is_same_v<Beta, dpf::bit>)
|
||||
return dpf::bit{static_cast<bool>(a) ^ static_cast<bool>(b)};
|
||||
else if constexpr (dpf::utils::is_xor_wrapper_v<Beta>)
|
||||
return Beta{static_cast<uint64_t>(a) ^ static_cast<uint64_t>(b)};
|
||||
return a ^ b;
|
||||
else
|
||||
return static_cast<Beta>(a - b);
|
||||
}
|
||||
|
|
@ -504,42 +505,50 @@ struct cmp_at_pack
|
|||
: if_true{std::move(t)}, if_false{std::move(f)} { }
|
||||
};
|
||||
|
||||
/// \complexity O(1). Builds a spec object. The n-level walk is `make_dpf` or `ds_advance_level`, not this function.
|
||||
template <typename Beta>
|
||||
inline auto lt(Beta t, Beta f = detail::dcf_impl::default_false<std::decay_t<Beta>>())
|
||||
{
|
||||
return cmp_pack<cmp_kind::lt, std::decay_t<Beta>>(std::move(t), std::move(f));
|
||||
}
|
||||
/// \complexity O(1). Builds a spec object. The n-level walk is `make_dpf` or `ds_advance_level`, not this function.
|
||||
template <typename Beta>
|
||||
inline auto leq(Beta t, Beta f = detail::dcf_impl::default_false<std::decay_t<Beta>>())
|
||||
{
|
||||
return cmp_pack<cmp_kind::leq, std::decay_t<Beta>>(std::move(t), std::move(f));
|
||||
}
|
||||
/// \complexity O(1). Builds a spec object. The n-level walk is `make_dpf` or `ds_advance_level`, not this function.
|
||||
template <typename Beta>
|
||||
inline auto gt(Beta t, Beta f = detail::dcf_impl::default_false<std::decay_t<Beta>>())
|
||||
{
|
||||
return cmp_pack<cmp_kind::gt, std::decay_t<Beta>>(std::move(t), std::move(f));
|
||||
}
|
||||
/// \complexity O(1). Builds a spec object. The n-level walk is `make_dpf` or `ds_advance_level`, not this function.
|
||||
template <typename Beta>
|
||||
inline auto geq(Beta t, Beta f = detail::dcf_impl::default_false<std::decay_t<Beta>>())
|
||||
{
|
||||
return cmp_pack<cmp_kind::geq, std::decay_t<Beta>>(std::move(t), std::move(f));
|
||||
}
|
||||
|
||||
/// \complexity O(1). Builds a spec object. The n-level walk is `make_dpf` or `ds_advance_level`, not this function.
|
||||
template <std::size_t N, typename Beta>
|
||||
inline auto lt_at(Beta t, Beta f = detail::dcf_impl::default_false<std::decay_t<Beta>>())
|
||||
{
|
||||
return cmp_at_pack<N, cmp_kind::lt, std::decay_t<Beta>>(std::move(t), std::move(f));
|
||||
}
|
||||
/// \complexity O(1). Builds a spec object. The n-level walk is `make_dpf` or `ds_advance_level`, not this function.
|
||||
template <std::size_t N, typename Beta>
|
||||
inline auto leq_at(Beta t, Beta f = detail::dcf_impl::default_false<std::decay_t<Beta>>())
|
||||
{
|
||||
return cmp_at_pack<N, cmp_kind::leq, std::decay_t<Beta>>(std::move(t), std::move(f));
|
||||
}
|
||||
/// \complexity O(1). Builds a spec object. The n-level walk is `make_dpf` or `ds_advance_level`, not this function.
|
||||
template <std::size_t N, typename Beta>
|
||||
inline auto gt_at(Beta t, Beta f = detail::dcf_impl::default_false<std::decay_t<Beta>>())
|
||||
{
|
||||
return cmp_at_pack<N, cmp_kind::gt, std::decay_t<Beta>>(std::move(t), std::move(f));
|
||||
}
|
||||
/// \complexity O(1). Builds a spec object. The n-level walk is `make_dpf` or `ds_advance_level`, not this function.
|
||||
template <std::size_t N, typename Beta>
|
||||
inline auto geq_at(Beta t, Beta f = detail::dcf_impl::default_false<std::decay_t<Beta>>())
|
||||
{
|
||||
|
|
@ -758,6 +767,7 @@ struct idcf_pack
|
|||
static constexpr std::size_t prefix = Spec::prefix;
|
||||
static constexpr std::size_t block_width = Spec::block_width;
|
||||
static constexpr std::size_t length_bits = spec_length_bits<Spec>::value;
|
||||
using inner_spec = Spec;
|
||||
using beta_type = typename Spec::beta_type;
|
||||
beta_type if_true;
|
||||
beta_type if_false;
|
||||
|
|
@ -817,6 +827,12 @@ inline auto eq_at(Beta t, Beta f = detail::dcf_impl::default_false<std::decay_t<
|
|||
return eq_at_pack<N, std::decay_t<Beta>>(std::move(t), std::move(f));
|
||||
}
|
||||
|
||||
template <typename T, typename = void>
|
||||
struct pack_has_fn : std::false_type {};
|
||||
template <typename T>
|
||||
struct pack_has_fn<T, std::void_t<decltype(std::declval<T &>().fn)>>
|
||||
: std::true_type {};
|
||||
|
||||
template <std::size_t BlockWidth, typename Spec>
|
||||
struct block_width_pack
|
||||
{
|
||||
|
|
@ -828,6 +844,7 @@ struct block_width_pack
|
|||
static constexpr std::size_t prefix = Spec::prefix;
|
||||
static constexpr bool incremental = spec_is_incremental<Spec>::value;
|
||||
static constexpr std::size_t length_bits = spec_length_bits<Spec>::value;
|
||||
using inner_spec = Spec;
|
||||
using beta_type = typename Spec::beta_type;
|
||||
beta_type if_true;
|
||||
beta_type if_false;
|
||||
|
|
@ -849,6 +866,99 @@ struct block_width_fn
|
|||
template <std::size_t BlockWidth>
|
||||
inline constexpr block_width_fn<BlockWidth> block_width{};
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Same comparison shape, wildcard payload. Wrappers (`idcf`, `block_width`,
|
||||
// `*_at`, path paints) keep their outer type; only the leaf beta changes.
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
template <typename T> struct is_idcf_pack : std::false_type {};
|
||||
template <typename Spec>
|
||||
struct is_idcf_pack<idcf_pack<Spec>> : std::true_type {};
|
||||
template <typename T>
|
||||
inline constexpr bool is_idcf_pack_v = is_idcf_pack<std::decay_t<T>>::value;
|
||||
|
||||
template <typename T> struct is_block_width_pack : std::false_type {};
|
||||
template <std::size_t B, typename Spec>
|
||||
struct is_block_width_pack<block_width_pack<B, Spec>> : std::true_type {};
|
||||
template <typename T>
|
||||
inline constexpr bool is_block_width_pack_v =
|
||||
is_block_width_pack<std::decay_t<T>>::value;
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
template <typename Spec>
|
||||
auto cmp_wildcard_shape();
|
||||
|
||||
template <cmp_kind K, typename B>
|
||||
auto cmp_wildcard_shape(cmp_pack<K, B>)
|
||||
{
|
||||
using W = wildcard_value<concrete_type_t<B>>;
|
||||
return cmp_pack<K, W>(W{}, W{});
|
||||
}
|
||||
template <std::size_t N, cmp_kind K, typename B>
|
||||
auto cmp_wildcard_shape(cmp_at_pack<N, K, B>)
|
||||
{
|
||||
using W = wildcard_value<concrete_type_t<B>>;
|
||||
return cmp_at_pack<N, K, W>(W{}, W{});
|
||||
}
|
||||
template <cmp_kind K, typename B, std::size_t L>
|
||||
auto cmp_wildcard_shape(paint_pack<K, B, L>)
|
||||
{
|
||||
using W = wildcard_value<concrete_type_t<B>>;
|
||||
return paint_pack<K, W, L>(W{}, W{});
|
||||
}
|
||||
template <std::size_t N, cmp_kind K, typename B, std::size_t L>
|
||||
auto cmp_wildcard_shape(paint_at_pack<N, K, B, L>)
|
||||
{
|
||||
using W = wildcard_value<concrete_type_t<B>>;
|
||||
return paint_at_pack<N, K, W, L>(W{}, W{});
|
||||
}
|
||||
template <typename B, typename Fn>
|
||||
auto cmp_wildcard_shape(paint_fn_pack<B, Fn> spec)
|
||||
{
|
||||
using W = wildcard_value<concrete_type_t<B>>;
|
||||
return paint_fn_pack<W, std::decay_t<Fn>>(W{}, W{}, std::move(spec.fn));
|
||||
}
|
||||
template <std::size_t N, typename B, typename Fn>
|
||||
auto cmp_wildcard_shape(paint_fn_at_pack<N, B, Fn> spec)
|
||||
{
|
||||
using W = wildcard_value<concrete_type_t<B>>;
|
||||
return paint_fn_at_pack<N, W, std::decay_t<Fn>>(W{}, W{}, std::move(spec.fn));
|
||||
}
|
||||
|
||||
template <typename Spec>
|
||||
auto cmp_wildcard_shape()
|
||||
{
|
||||
if constexpr (is_idcf_pack_v<Spec>)
|
||||
return idcf(cmp_wildcard_shape<typename Spec::inner_spec>());
|
||||
else if constexpr (is_block_width_pack_v<Spec>)
|
||||
return block_width<Spec::block_width>(
|
||||
cmp_wildcard_shape<typename Spec::inner_spec>());
|
||||
else
|
||||
{
|
||||
using B = typename Spec::beta_type;
|
||||
return cmp_wildcard_shape(Spec{B{}, B{}});
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/// @brief The same comparison or paint spec, with a wildcard payload.
|
||||
/// @details `idcf`, `block_width`, and `*_at` stay wrapped. A spec that is
|
||||
/// already wildcard is returned unchanged. `path_paint` keeps `fn`.
|
||||
template <typename Spec>
|
||||
auto cmp_spec_as_wildcard(Spec spec)
|
||||
{
|
||||
using S = std::decay_t<Spec>;
|
||||
if constexpr (is_wildcard_v<typename S::beta_type>)
|
||||
return spec;
|
||||
else if constexpr (pack_has_fn<S>::value)
|
||||
return detail::cmp_wildcard_shape(std::move(spec));
|
||||
else
|
||||
return detail::cmp_wildcard_shape<S>();
|
||||
}
|
||||
|
||||
template <typename T> struct is_cmp_spec : std::false_type {};
|
||||
template <cmp_kind K, typename B> struct is_cmp_spec<cmp_pack<K, B>> : std::true_type {};
|
||||
template <std::size_t N, cmp_kind K, typename B>
|
||||
|
|
|
|||
303
include/dpf/deferred_rotated_subinterval.hpp
Normal file
303
include/dpf/deferred_rotated_subinterval.hpp
Normal file
|
|
@ -0,0 +1,303 @@
|
|||
/// @file dpf/deferred_rotated_subinterval.hpp
|
||||
/// @brief Deferred view over a full-domain buffer until the input offset is set.
|
||||
/// @details After `defer_eval_interval` / `defer_eval_full` fills a full-domain
|
||||
/// buffer at identity, this view waits for `assign_wildcard_input`.
|
||||
/// Materialization builds `rotation_iterable` by `offset_x(0)`. The
|
||||
/// logical `[from, to]` is then indexed in domain-walk order (same as
|
||||
/// `eval_full`), including wrap-around when `from > to` in that order
|
||||
/// (e.g. unsigned `[200, 10]`).
|
||||
/// @see dpf::defer_eval_interval
|
||||
/// @see dpf::rotation_iterable
|
||||
/// @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_DPF_DEFERRED_ROTATED_SUBINTERVAL_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_DEFERRED_ROTATED_SUBINTERVAL_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <iterator>
|
||||
#include <limits>
|
||||
#include <optional>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/rotation_iterable.hpp"
|
||||
#include "dpf/utils.hpp"
|
||||
#include "dpf/wildcard.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
/// @brief Index-based range over a (possibly wrapping) slice of a rotation.
|
||||
/// @details Indexing uses `rotation_iterable::operator[]`, so wrapping
|
||||
/// intervals stay correct without requiring a contiguous iterator
|
||||
/// walk past `end()`.
|
||||
template <typename Rotation>
|
||||
class deferred_rotated_range
|
||||
{
|
||||
public:
|
||||
class iterator
|
||||
{
|
||||
public:
|
||||
using iterator_category = std::bidirectional_iterator_tag;
|
||||
using difference_type = std::ptrdiff_t;
|
||||
using value_type = std::decay_t<decltype(std::declval<const Rotation &>()[0])>;
|
||||
using reference = decltype(std::declval<const Rotation &>()[0]);
|
||||
using pointer = void;
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr iterator() noexcept
|
||||
: rot_{nullptr}, start_{0}, pos_{0}, count_{0}, n_{0}
|
||||
{ }
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr iterator(const Rotation * rot, std::size_t start,
|
||||
std::size_t pos, std::size_t count, std::size_t n) noexcept
|
||||
: rot_{rot}, start_{start}, pos_{pos}, count_{count}, n_{n}
|
||||
{ }
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
reference operator*() const
|
||||
{
|
||||
const std::size_t idx = (start_ + pos_) % n_;
|
||||
return (*rot_)[static_cast<typename Rotation::difference_type>(idx)];
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr iterator & operator++() noexcept
|
||||
{
|
||||
++pos_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr iterator operator++(int) noexcept
|
||||
{
|
||||
iterator tmp = *this;
|
||||
++(*this);
|
||||
return tmp;
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr iterator & operator--() noexcept
|
||||
{
|
||||
--pos_;
|
||||
return *this;
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr iterator operator--(int) noexcept
|
||||
{
|
||||
iterator tmp = *this;
|
||||
--(*this);
|
||||
return tmp;
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr bool operator==(const iterator & rhs) const noexcept
|
||||
{
|
||||
return pos_ == rhs.pos_ && rot_ == rhs.rot_
|
||||
&& start_ == rhs.start_ && count_ == rhs.count_
|
||||
&& n_ == rhs.n_;
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr bool operator!=(const iterator & rhs) const noexcept
|
||||
{
|
||||
return !(*this == rhs);
|
||||
}
|
||||
|
||||
private:
|
||||
const Rotation * rot_;
|
||||
std::size_t start_;
|
||||
std::size_t pos_;
|
||||
std::size_t count_;
|
||||
std::size_t n_;
|
||||
};
|
||||
|
||||
using const_iterator = iterator;
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr deferred_rotated_range(const Rotation * rot, std::size_t start,
|
||||
std::size_t count, std::size_t n) noexcept
|
||||
: rot_{rot}, start_{start}, count_{count}, n_{n}
|
||||
{ }
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr iterator begin() const noexcept
|
||||
{
|
||||
return iterator(rot_, start_, 0, count_, n_);
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr iterator end() const noexcept
|
||||
{
|
||||
return iterator(rot_, start_, count_, count_, n_);
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr const_iterator cbegin() const noexcept
|
||||
{
|
||||
return begin();
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr const_iterator cend() const noexcept
|
||||
{
|
||||
return end();
|
||||
}
|
||||
|
||||
private:
|
||||
const Rotation * rot_;
|
||||
std::size_t start_;
|
||||
std::size_t count_;
|
||||
std::size_t n_;
|
||||
};
|
||||
|
||||
/// @brief Full-domain buffer view that rotates after the input offset is ready.
|
||||
/// @tparam DpfKey DPF key type (wildcard input)
|
||||
/// @tparam IteratorT random-access iterator into the full-domain output buffer
|
||||
template <typename DpfKey,
|
||||
typename IteratorT>
|
||||
class deferred_rotated_subinterval
|
||||
{
|
||||
public:
|
||||
using dpf_type = DpfKey;
|
||||
using input_type = typename dpf_type::input_type;
|
||||
using iterator = IteratorT;
|
||||
using size_type = std::size_t;
|
||||
using difference_type =
|
||||
typename std::iterator_traits<iterator>::difference_type;
|
||||
using rotation_type = rotation_iterable<iterator>;
|
||||
using view_type = deferred_rotated_range<rotation_type>;
|
||||
|
||||
/// @param dpf key whose `offset_x` will supply the rotation once ready
|
||||
/// @param begin begin of the full-domain buffer (identity eval order)
|
||||
/// @param end end of the full-domain buffer
|
||||
/// @param from inclusive logical start
|
||||
/// @param to inclusive logical end
|
||||
/// @param outputs_per_leaf leaf packing width (usually `DpfKey::outputs_per_leaf`)
|
||||
HEDLEY_NO_THROW
|
||||
deferred_rotated_subinterval(const dpf_type & dpf, iterator begin,
|
||||
iterator end, input_type from, input_type to,
|
||||
size_type outputs_per_leaf) noexcept
|
||||
: dpf_{dpf},
|
||||
begin_{begin},
|
||||
end_{end},
|
||||
from_{from},
|
||||
to_{to},
|
||||
outputs_{outputs_per_leaf},
|
||||
rot_{std::nullopt}
|
||||
{
|
||||
(void)outputs_;
|
||||
}
|
||||
|
||||
/// @brief Materialize the rotated subinterval; requires assigned input.
|
||||
/// @return a bidirectional range over the logical `[from, to]`
|
||||
/// @note The returned range borrows this object's cached rotation; keep
|
||||
/// `*this` alive for the range's lifetime.
|
||||
view_type get()
|
||||
{
|
||||
ensure_rotation();
|
||||
return make_view();
|
||||
}
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto begin()
|
||||
{
|
||||
return get().begin();
|
||||
}
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto end()
|
||||
{
|
||||
return get().end();
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr const dpf_type & dpf() const noexcept
|
||||
{
|
||||
return dpf_;
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr input_type from() const noexcept
|
||||
{
|
||||
return from_;
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr input_type to() const noexcept
|
||||
{
|
||||
return to_;
|
||||
}
|
||||
|
||||
private:
|
||||
static constexpr auto to_integral_t =
|
||||
utils::to_integral_type<input_type>{};
|
||||
static constexpr auto bits = utils::bitlength_of_v<input_type>;
|
||||
|
||||
void ensure_rotation()
|
||||
{
|
||||
assert_not_wildcard_input(dpf_);
|
||||
if (!rot_.has_value())
|
||||
{
|
||||
const auto offset = dpf_.offset_x(input_type{});
|
||||
rot_ = rotation_type(begin_, end_,
|
||||
static_cast<difference_type>(to_integral_t(offset)));
|
||||
}
|
||||
}
|
||||
|
||||
view_type make_view()
|
||||
{
|
||||
// Domain walk of `eval_full`: index 0 is `min`, then `++` order.
|
||||
const size_type n =
|
||||
static_cast<size_type>(std::distance(begin_, end_));
|
||||
const auto min_i =
|
||||
to_integral_t(std::numeric_limits<input_type>::min());
|
||||
auto from_i = to_integral_t(from_);
|
||||
auto span = to_integral_t(to_) - from_i;
|
||||
if constexpr (bits < utils::bitlength_of_v<decltype(span)>)
|
||||
{
|
||||
span &= (decltype(span){1} << bits) - 1;
|
||||
}
|
||||
auto from_rel = from_i - min_i;
|
||||
if constexpr (bits < utils::bitlength_of_v<decltype(from_rel)>)
|
||||
{
|
||||
from_rel &= (decltype(from_rel){1} << bits) - 1;
|
||||
}
|
||||
const auto start = static_cast<size_type>(from_rel);
|
||||
const auto count = static_cast<size_type>(span) + size_type{1};
|
||||
return view_type(&rot_.value(), start, count, n);
|
||||
}
|
||||
|
||||
const dpf_type & dpf_;
|
||||
iterator begin_;
|
||||
iterator end_;
|
||||
input_type from_;
|
||||
input_type to_;
|
||||
size_type outputs_;
|
||||
std::optional<rotation_type> rot_;
|
||||
};
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_DEFERRED_ROTATED_SUBINTERVAL_HPP__
|
||||
|
|
@ -1,12 +1,16 @@
|
|||
/// @file dpf/doerner_shelat.hpp
|
||||
/// @brief Doerner–Shelat generation of a dealer DPF key.
|
||||
/// @details Two shares of the point are walked level by level — XOR shares by
|
||||
/// default, or additive shares when tagged with `arith_input`.
|
||||
/// default. `arith_input` holds additive shares of the point. A beaver
|
||||
/// ripple-carry converts them to XOR shares of the sum bits, and that
|
||||
/// sharing is what the walk consumes. The sum is not opened.
|
||||
/// Correction words, advice bits, seeds, and leaves are the ones
|
||||
/// `make_dpf` would emit for the reconstructed point, the same roots,
|
||||
/// and the same beaver coins. Beaver pads used to hide the path bit
|
||||
/// cancel and are not part of the key. Pad randomness must not come
|
||||
/// from `uniform_fill` if the beaver tape is being matched.
|
||||
/// `make_dpf` would emit for that point, the same roots, and the same
|
||||
/// beaver coins. Beaver pads used to hide a path bit cancel and are
|
||||
/// not part of the key. Pad randomness must not come from
|
||||
/// `uniform_fill` if the beaver tape is being matched.
|
||||
/// @note Following Jack Doerner and abhi shelat, CCS 2017 (ePrint 2017/827): one correction word opened per level from shares of the point.
|
||||
/// @note Guo, Yang, Wang, Zhang, Xie, Zhang, and Liu (ePrint 2022/1431, §5.2) generate a DPF in the COT/OLE hybrid in n+3 rounds, with no beaver-pad dealer. This header uses that dealer tape and one opening round per level.
|
||||
/// @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.
|
||||
|
|
@ -16,16 +20,23 @@
|
|||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
#include "simde/simde/x86/avx2.h"
|
||||
|
||||
#include "dpf/dpf_key.hpp"
|
||||
#include "dpf/experiment_note.hpp"
|
||||
#include "dpf/random.hpp"
|
||||
#include "dpf/dcf.hpp"
|
||||
#include "dpf/constrained_cmp.hpp"
|
||||
#include "dpf/beaver.hpp"
|
||||
#include "dpf/xor_wrapper.hpp"
|
||||
#include "dpf/leaf_node.hpp"
|
||||
#include "dpf/leaf_arithmetic.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
|
@ -99,6 +110,40 @@ struct ds_randomness
|
|||
PadRng pad;
|
||||
};
|
||||
|
||||
/// @brief Pad stream whose `block()` and `bit()` come from one PRG seed.
|
||||
/// @details Drop-in for `detail::urandom_pad_rng` on Doerner–Shelat dealers.
|
||||
/// `pseudorandom_root_sampler<PRG>` is the matching root source.
|
||||
template <typename PRG = dpf::prg::aes128>
|
||||
struct prg_pad_rng
|
||||
{
|
||||
using block_type = typename PRG::block_type;
|
||||
|
||||
explicit prg_pad_rng(block_type seed = dpf::uniform_sample<block_type>())
|
||||
: seed_(seed)
|
||||
{
|
||||
note_experiment_seed("prg_pad_rng", seed_);
|
||||
}
|
||||
|
||||
block_type block()
|
||||
{
|
||||
return PRG::eval(seed_, n_++);
|
||||
}
|
||||
|
||||
uint8_t bit()
|
||||
{
|
||||
const block_type drawn = block();
|
||||
unsigned char low = 0;
|
||||
std::memcpy(&low, &drawn, 1);
|
||||
return static_cast<uint8_t>(low & 1u);
|
||||
}
|
||||
|
||||
const block_type & seed() const noexcept { return seed_; }
|
||||
|
||||
private:
|
||||
block_type seed_{};
|
||||
std::uint32_t n_ = 0;
|
||||
};
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
|
|
@ -164,31 +209,128 @@ simde__m128i ds_gate(uint8_t bit, simde__m128i block) noexcept
|
|||
template <typename PadRng>
|
||||
ds_cw_pads ds_sample_cw(PadRng & pad)
|
||||
{
|
||||
// Ideal (semi-honest): each party holds (rand, bit, gamma) with
|
||||
// gamma0 ⊕ gamma1 = (bit1 · rand0) ⊕ (bit0 · rand1).
|
||||
// The full product (bit1 · rand0) is NOT given to party 0: that would
|
||||
// leak bit1, and with the opened blind bit = path1 ⊕ bit1 it would
|
||||
// open the peer path bit (and thus α under an oblivious walk). Shares
|
||||
// of the XOR of both products hide both pad bits from each party.
|
||||
// The product share is a beaver session over the XOR ring; the clear
|
||||
// bit and block stay with the party that owns them.
|
||||
using Ring = dpf::xor_wrapper<simde_uint128>;
|
||||
using traits = dpf::beavers::ring_traits<Ring>;
|
||||
auto to_ring = [](simde__m128i block) {
|
||||
simde_uint128 raw{};
|
||||
std::memcpy(&raw, &block, sizeof(block));
|
||||
return Ring{raw};
|
||||
};
|
||||
auto to_block = [](const Ring & ring) {
|
||||
simde__m128i block{};
|
||||
const auto raw = static_cast<typename Ring::value_type>(ring);
|
||||
std::memcpy(&block, &raw, sizeof(block));
|
||||
return block;
|
||||
};
|
||||
const Ring rand0 = to_ring(pad.block());
|
||||
const Ring rand1 = to_ring(pad.block());
|
||||
const bool bit0 = (pad.bit() & 1u) != 0;
|
||||
const bool bit1 = (pad.bit() & 1u) != 0;
|
||||
auto sampler = [&pad, &to_ring]() -> Ring {
|
||||
return to_ring(pad.block());
|
||||
};
|
||||
|
||||
dpf::beavers::session<Ring> s;
|
||||
auto b0 = s.bit();
|
||||
auto b1 = s.bit();
|
||||
auto r0 = s.input();
|
||||
auto r1 = s.input();
|
||||
auto prod = s(b1 * r0 + b0 * r1);
|
||||
s.pin(prod);
|
||||
s.sample(sampler);
|
||||
s.bind(b0, bit0 ? traits::one() : traits::zero(), sampler);
|
||||
s.bind(b1, bit1 ? traits::one() : traits::zero(), sampler);
|
||||
s.bind(r0, rand0, sampler);
|
||||
s.bind(r1, rand1, sampler);
|
||||
s.evaluate();
|
||||
const auto gamma = s.value(prod);
|
||||
|
||||
ds_cw_pads p{};
|
||||
const simde__m128i zero = simde_mm_setzero_si128();
|
||||
p.p0.rand = pad.block();
|
||||
p.p1.rand = pad.block();
|
||||
p.p0.bit = static_cast<uint8_t>(pad.bit() & 1u);
|
||||
p.p1.bit = static_cast<uint8_t>(pad.bit() & 1u);
|
||||
p.p0.gamma = p.p1.bit ? p.p0.rand : zero;
|
||||
p.p1.gamma = p.p0.bit ? p.p1.rand : zero;
|
||||
p.p0.rand = to_block(rand0);
|
||||
p.p1.rand = to_block(rand1);
|
||||
p.p0.bit = static_cast<uint8_t>(bit0);
|
||||
p.p1.bit = static_cast<uint8_t>(bit1);
|
||||
p.p0.gamma = to_block(gamma.p0);
|
||||
p.p1.gamma = to_block(gamma.p1);
|
||||
return p;
|
||||
}
|
||||
|
||||
/// @brief Pack a XOR-ring `bit_mul` into the classical DS AND pad shape.
|
||||
/// @tparam Ring XOR ring whose unit is the all-ones word
|
||||
/// @param bm the sampled bit-mul material
|
||||
/// @param a0_share party 0's XOR share of the opened bit
|
||||
/// @return pads ready for `ds_and_open`
|
||||
template <typename Ring>
|
||||
ds_and_pads ds_and_from_bit_mul(const dpf::beavers::bit_mul_beaver<Ring> & bm,
|
||||
uint8_t a0_share)
|
||||
{
|
||||
using traits = dpf::beavers::ring_traits<Ring>;
|
||||
const Ring opened = bm.bit.open();
|
||||
const uint8_t a = (opened == traits::one()) ? uint8_t{1} : uint8_t{0};
|
||||
ds_and_pads p{};
|
||||
p.a0 = static_cast<uint8_t>(a0_share & 1u);
|
||||
p.a1 = static_cast<uint8_t>(a ^ p.a0);
|
||||
auto to_block = [](const Ring & r) {
|
||||
simde__m128i b{};
|
||||
const auto v = static_cast<typename Ring::value_type>(r);
|
||||
static_assert(sizeof(v) == sizeof(simde__m128i),
|
||||
"ds AND packs a 128-bit XOR ring into an AES block");
|
||||
std::memcpy(&b, &v, sizeof(b));
|
||||
return b;
|
||||
};
|
||||
p.b0_share = to_block(bm.scalar.p0);
|
||||
p.b1_share = to_block(bm.scalar.p1);
|
||||
p.c0_share = to_block(bm.product.p0);
|
||||
p.c1_share = to_block(bm.product.p1);
|
||||
return p;
|
||||
}
|
||||
|
||||
/// @brief Sample one DS AND from `sample_bit_mul`, driven by `pad` or `rng`.
|
||||
/// @tparam PadRng pad stream with `block()` / `bit()`
|
||||
/// @tparam Sample callable returning a 128-bit XOR ring element
|
||||
/// @param pad the Doerner–Shelat pad stream (bit for the clear Beaver bit)
|
||||
/// @param rng ring sampler for `sample_bit_mul` (defaults to `pad.block()`)
|
||||
/// @return classical AND pads for `ds_and_open`
|
||||
template <typename PadRng, typename Sample>
|
||||
ds_and_pads ds_sample_and(PadRng & pad, Sample && rng)
|
||||
{
|
||||
using Ring = dpf::xor_wrapper<simde_uint128>;
|
||||
using traits = dpf::beavers::ring_traits<Ring>;
|
||||
auto & ring_rng = rng;
|
||||
std::size_t phase = 0;
|
||||
auto sampler = [&]() -> Ring {
|
||||
if (phase == 0)
|
||||
{
|
||||
++phase;
|
||||
return (pad.bit() & 1u) ? traits::one() : traits::zero();
|
||||
}
|
||||
++phase;
|
||||
return ring_rng();
|
||||
};
|
||||
auto bm = dpf::beavers::sample_bit_mul<Ring>(sampler);
|
||||
const uint8_t a0 = static_cast<uint8_t>(pad.bit() & 1u);
|
||||
return ds_and_from_bit_mul(bm, a0);
|
||||
}
|
||||
|
||||
template <typename PadRng>
|
||||
ds_and_pads ds_sample_and(PadRng & pad)
|
||||
{
|
||||
ds_and_pads p{};
|
||||
const uint8_t a = static_cast<uint8_t>(pad.bit() & 1u);
|
||||
const simde__m128i B = pad.block();
|
||||
const simde__m128i C = ds_gate(a, B);
|
||||
p.a0 = static_cast<uint8_t>(pad.bit() & 1u);
|
||||
p.a1 = static_cast<uint8_t>(a ^ p.a0);
|
||||
p.b0_share = pad.block();
|
||||
p.b1_share = ds_xor(B, p.b0_share);
|
||||
p.c0_share = pad.block();
|
||||
p.c1_share = ds_xor(C, p.c0_share);
|
||||
return p;
|
||||
using Ring = dpf::xor_wrapper<simde_uint128>;
|
||||
auto from_pad = [&pad]() -> Ring {
|
||||
const simde__m128i b = pad.block();
|
||||
simde_uint128 v{};
|
||||
std::memcpy(&v, &b, sizeof(b));
|
||||
return Ring{v};
|
||||
};
|
||||
return ds_sample_and(pad, from_pad);
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
|
|
@ -279,6 +421,123 @@ inline simde__m128i ds_deliver(uint8_t b_exp, simde__m128i base, simde__m128i M,
|
|||
return ds_xor(ds_xor(local, z.z0), z.z1);
|
||||
}
|
||||
|
||||
/// @brief XOR shares of a bit-Beaver triple `(α, β, α∧β)`.
|
||||
struct ds_bit_triple
|
||||
{
|
||||
uint8_t a0;
|
||||
uint8_t a1;
|
||||
uint8_t b0;
|
||||
uint8_t b1;
|
||||
uint8_t c0;
|
||||
uint8_t c1;
|
||||
};
|
||||
|
||||
/// @brief Sample one bit-AND triple from `pad`, via `sample_beaver2` on the XOR ring.
|
||||
/// @tparam PadRng pad stream with `bit()`
|
||||
/// @param pad the pad stream
|
||||
/// @return shares of `(α, β, α∧β)`
|
||||
template <typename PadRng>
|
||||
ds_bit_triple ds_sample_bit_and(PadRng & pad)
|
||||
{
|
||||
using Bit = dpf::xor_wrapper<std::uint8_t>;
|
||||
auto sampler = [&pad]() -> Bit {
|
||||
std::uint8_t packed = 0;
|
||||
for (int i = 0; i < 8; ++i)
|
||||
{
|
||||
packed = static_cast<std::uint8_t>(
|
||||
packed | (static_cast<std::uint8_t>(pad.bit() & 1u) << i));
|
||||
}
|
||||
return Bit{packed};
|
||||
};
|
||||
const auto triple = dpf::beavers::sample_beaver2<Bit>(sampler);
|
||||
auto low = [](Bit x) {
|
||||
return static_cast<uint8_t>(static_cast<std::uint8_t>(x) & 1u);
|
||||
};
|
||||
return ds_bit_triple{
|
||||
low(triple.a.p0), low(triple.a.p1),
|
||||
low(triple.b.p0), low(triple.b.p1),
|
||||
low(triple.ab.p0), low(triple.ab.p1)};
|
||||
}
|
||||
|
||||
/// @brief One party's share of `x ∧ y` after `d = x⊕α` and `e = y⊕β` are open.
|
||||
/// @param d the opened mask of `x`
|
||||
/// @param e the opened mask of `y`
|
||||
/// @param a this party's share of `α`
|
||||
/// @param b this party's share of `β`
|
||||
/// @param c this party's share of `α∧β`
|
||||
/// @param hold_de party 0 adds the public `d∧e` term
|
||||
/// @return this party's XOR share of the product
|
||||
HEDLEY_NO_THROW
|
||||
inline uint8_t ds_bit_and_party(uint8_t d, uint8_t e, uint8_t a, uint8_t b,
|
||||
uint8_t c, bool hold_de) noexcept
|
||||
{
|
||||
uint8_t z = static_cast<uint8_t>((d & b) ^ (e & a) ^ c);
|
||||
if (hold_de)
|
||||
z = static_cast<uint8_t>(z ^ (d & e));
|
||||
return static_cast<uint8_t>(z & 1u);
|
||||
}
|
||||
|
||||
/// @brief Joint evaluation of one bit-AND. `d` and `e` are the opened masks.
|
||||
/// @param t the triple
|
||||
/// @param x0 party 0's share of `x`
|
||||
/// @param x1 party 1's share of `x`
|
||||
/// @param y0 party 0's share of `y`
|
||||
/// @param y1 party 1's share of `y`
|
||||
/// @return XOR shares of `x ∧ y`
|
||||
HEDLEY_NO_THROW
|
||||
inline std::pair<uint8_t, uint8_t> ds_bit_and_shares(const ds_bit_triple & t,
|
||||
uint8_t x0, uint8_t x1, uint8_t y0, uint8_t y1) noexcept
|
||||
{
|
||||
const uint8_t d = static_cast<uint8_t>(x0 ^ x1 ^ t.a0 ^ t.a1);
|
||||
const uint8_t e = static_cast<uint8_t>(y0 ^ y1 ^ t.b0 ^ t.b1);
|
||||
return {ds_bit_and_party(d, e, t.a0, t.b0, t.c0, true),
|
||||
ds_bit_and_party(d, e, t.a1, t.b1, t.c1, false)};
|
||||
}
|
||||
|
||||
/// @brief Replace additive shares with XOR shares of their sum.
|
||||
/// @details One beaver bit-AND per bit except the last. Party 0's sum-bit
|
||||
/// share is `a ⊕ c0`; party 1's is `b ⊕ c1`. The carry share is
|
||||
/// `((a⊕c) ∧ (b⊕c)) ⊕ c`, which is the majority. Neither share is
|
||||
/// the sum, and the sum is not written down.
|
||||
/// @tparam PadRng pad stream with `bit()`
|
||||
/// @tparam InputT input domain type
|
||||
/// @param pads the pad stream
|
||||
/// @param x0 party 0's additive share, replaced by its XOR share of the sum
|
||||
/// @param x1 party 1's additive share, replaced by its XOR share of the sum
|
||||
template <typename PadRng, typename InputT>
|
||||
void split_additive_to_xor(PadRng & pads, InputT & x0, InputT & x1)
|
||||
{
|
||||
constexpr auto to_int = utils::to_integral_type<InputT>{};
|
||||
using FromI = typename utils::make_from_integral_value<InputT>::integral_type;
|
||||
using U = std::make_unsigned_t<FromI>;
|
||||
const U u0 = static_cast<U>(to_int(x0));
|
||||
const U u1 = static_cast<U>(to_int(x1));
|
||||
U s0 = 0;
|
||||
U s1 = 0;
|
||||
uint8_t c0 = 0;
|
||||
uint8_t c1 = 0;
|
||||
constexpr std::size_t nbits = utils::bitlength_of_v<InputT>;
|
||||
for (std::size_t i = 0; i < nbits; ++i)
|
||||
{
|
||||
const uint8_t a = static_cast<uint8_t>((u0 >> i) & U{1});
|
||||
const uint8_t b = static_cast<uint8_t>((u1 >> i) & U{1});
|
||||
const uint8_t sum0 = static_cast<uint8_t>(a ^ c0);
|
||||
const uint8_t sum1 = static_cast<uint8_t>(b ^ c1);
|
||||
s0 = static_cast<U>(s0 | (static_cast<U>(sum0) << i));
|
||||
s1 = static_cast<U>(s1 | (static_cast<U>(sum1) << i));
|
||||
if (i + 1 == nbits)
|
||||
break;
|
||||
const ds_bit_triple triple = ds_sample_bit_and(pads);
|
||||
const auto prod = ds_bit_and_shares(triple,
|
||||
static_cast<uint8_t>(a ^ c0), c1,
|
||||
c0, static_cast<uint8_t>(b ^ c1));
|
||||
c0 = static_cast<uint8_t>(prod.first ^ c0);
|
||||
c1 = static_cast<uint8_t>(prod.second ^ c1);
|
||||
}
|
||||
x0 = utils::make_from_integral_value<InputT>{}(static_cast<FromI>(s0));
|
||||
x1 = utils::make_from_integral_value<InputT>{}(static_cast<FromI>(s1));
|
||||
}
|
||||
|
||||
/// @brief Per-level messages prepared before the CW protocol runs (blinds + pads).
|
||||
struct ds_level_blinds
|
||||
{
|
||||
|
|
@ -326,6 +585,35 @@ struct ds_cmp_gen_state
|
|||
const void * paint_ctx = nullptr;
|
||||
};
|
||||
|
||||
/// @brief Local joint-sim mux of a packed naked leaf from XOR bit shares.
|
||||
/// @details Mirrors `party/oblivious_select.hpp` `mux_leaf_share`: each level
|
||||
/// selects with `bit0 ^ bit1` so the call site never forms a clear
|
||||
/// point for `make_leaves`. The joint simulator holds both shares.
|
||||
/// @tparam Leaf packed leaf type
|
||||
/// @tparam Make candidate builder `Leaf(unsigned lane)`
|
||||
/// @tparam Bit0 party-0 bit accessor
|
||||
/// @tparam Bit1 party-1 bit accessor
|
||||
template <typename Leaf, typename Make, typename Bit0, typename Bit1>
|
||||
Leaf mux_naked_leaf_local(std::size_t lg, Make && make, Bit0 && bit0_at,
|
||||
Bit1 && bit1_at)
|
||||
{
|
||||
if (lg == 0)
|
||||
return make(0u);
|
||||
std::vector<Leaf> cand(std::size_t{1} << lg);
|
||||
for (std::size_t i = 0; i < cand.size(); ++i)
|
||||
cand[i] = make(static_cast<unsigned>(i));
|
||||
for (std::size_t b = 0; b < lg; ++b)
|
||||
{
|
||||
const uint8_t bit = static_cast<uint8_t>(
|
||||
(bit0_at(b) ^ bit1_at(b)) & 1u);
|
||||
std::vector<Leaf> next(cand.size() / 2);
|
||||
for (std::size_t k = 0; k < next.size(); ++k)
|
||||
next[k] = bit ? cand[2 * k + 1] : cand[2 * k];
|
||||
cand.swap(next);
|
||||
}
|
||||
return cand[0];
|
||||
}
|
||||
|
||||
/// @brief Local joint simulation: today's `ds_cw_outs` / `ds_open_advice` / `ds_and_open`.
|
||||
/// @details An MPC backend would send `blinds` and return the same `ds_level_open` shape.
|
||||
/// @tparam PadRng pad stream for the Doerner–Shelat protocol
|
||||
|
|
@ -451,83 +739,23 @@ struct local_cw_protocol
|
|||
std::forward<BlockSampler>(sample));
|
||||
}
|
||||
|
||||
/// @brief Majority of three bits (next carry of a full adder).
|
||||
/// @param a the `a`
|
||||
/// @param b the `b`
|
||||
/// @param c the `c`
|
||||
/// @return Majority of three bits (next carry of a full adder)
|
||||
HEDLEY_NO_THROW
|
||||
static constexpr uint8_t majority(uint8_t a, uint8_t b, uint8_t c) noexcept
|
||||
{
|
||||
return static_cast<uint8_t>((a & b) | (a & c) | (b & c));
|
||||
}
|
||||
|
||||
/// @brief One additive digit: sum bit `a XOR b XOR cin`, carry out = majority.
|
||||
/// @param a the `a`
|
||||
/// @param b the `b`
|
||||
/// @param cin the `cin`
|
||||
/// @param cout the `cout`
|
||||
/// @return One additive digit: sum bit `a XOR b XOR cin`, carry out = majority
|
||||
HEDLEY_NO_THROW
|
||||
static constexpr uint8_t open_sum_bit(uint8_t a, uint8_t b, uint8_t cin,
|
||||
uint8_t & cout) noexcept
|
||||
{
|
||||
cout = majority(a, b, cin);
|
||||
return static_cast<uint8_t>(a ^ b ^ cin);
|
||||
}
|
||||
|
||||
/// @brief Reconstruct `a0 + a1` via a LSB→MSB carry chain, then flip the MSB
|
||||
/// when the domain is signed — matching `make_dpf` on the sum. The call
|
||||
/// site never forms the sum; an MPC backend would open the same bits.
|
||||
/// @brief Encode shares for the XOR-style CW walk.
|
||||
/// @details XOR inputs flip party 0's MSB, which is linear over XOR.
|
||||
/// Additive inputs are converted first: `split_additive_to_xor`
|
||||
/// draws one bit-Beaver per carry and leaves XOR shares of the
|
||||
/// sum. The MSB flip is then the same one XOR inputs take, so
|
||||
/// the walk matches `make_dpf` on the sum. The sum is not opened
|
||||
/// and party 1's share is not cleared.
|
||||
/// @tparam InputT input domain type
|
||||
/// @param a0 the `a0`
|
||||
/// @param a1 the `a1`
|
||||
/// @return Reconstruct `a0 + a1` via a LSB→MSB carry chain, then flip the MSB when the domain
|
||||
/// is signed — matching `make_dpf` on the sum
|
||||
/// @param x0 party 0's share
|
||||
/// @param x1 party 1's share
|
||||
/// @param arith `true` when `x0`, `x1` are additive
|
||||
template <typename InputT>
|
||||
InputT open_arith_point(InputT a0, InputT a1) const
|
||||
{
|
||||
constexpr auto to_int = utils::to_integral_type<InputT>{};
|
||||
using FromI = typename utils::make_from_integral_value<InputT>::integral_type;
|
||||
using U = std::make_unsigned_t<FromI>;
|
||||
const U u0 = static_cast<U>(to_int(a0));
|
||||
const U u1 = static_cast<U>(to_int(a1));
|
||||
U sum = 0;
|
||||
uint8_t carry = 0;
|
||||
constexpr std::size_t nbits = utils::bitlength_of_v<InputT>;
|
||||
for (std::size_t i = 0; i < nbits; ++i)
|
||||
{
|
||||
const uint8_t b0 = static_cast<uint8_t>((u0 >> i) & U{1});
|
||||
const uint8_t b1 = static_cast<uint8_t>((u1 >> i) & U{1});
|
||||
const uint8_t s = open_sum_bit(b0, b1, carry, carry);
|
||||
sum = static_cast<U>(sum | (static_cast<U>(s) << i));
|
||||
}
|
||||
InputT out = utils::make_from_integral_value<InputT>{}(
|
||||
static_cast<FromI>(sum));
|
||||
utils::flip_msb_if_signed_integral(out);
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief Encode shares for the XOR-style CW walk. XOR mode flips party 0's MSB
|
||||
/// (linear over XOR). Arithmetic mode opens the sum (carry + signed MSB)
|
||||
/// and returns `(alpha, 0)` so the walk matches `make_dpf(alpha)`.
|
||||
/// @tparam InputT input domain type
|
||||
/// @param x0 the `x0`
|
||||
/// @param x1 the `x1`
|
||||
/// @param arith the `arith`
|
||||
template <typename InputT>
|
||||
void encode_walk_shares(InputT & x0, InputT & x1, bool arith) const
|
||||
void encode_walk_shares(InputT & x0, InputT & x1, bool arith)
|
||||
{
|
||||
if (arith)
|
||||
{
|
||||
const InputT alpha = open_arith_point(x0, x1);
|
||||
x0 = alpha;
|
||||
x1 = InputT{};
|
||||
}
|
||||
else
|
||||
{
|
||||
utils::flip_msb_if_signed_integral(x0);
|
||||
}
|
||||
split_additive_to_xor(pads, x0, x1);
|
||||
utils::flip_msb_if_signed_integral(x0);
|
||||
}
|
||||
|
||||
/// @brief Constrained comparison Π_CCMP: open `1{x0 < x1}` when `|x0−x1|=1`.
|
||||
|
|
@ -542,14 +770,18 @@ struct local_cw_protocol
|
|||
}
|
||||
|
||||
/// @brief Open a public leaf CW for a shared payload.
|
||||
/// @details Ring: `β = y0 + y1`; `g = CCMP(t0,t1)` selects `β − M` vs `M − β`
|
||||
/// (matches `make_leaf` with `sign = t0`). Characteristic 2: `β = y0 ⊕ y1`
|
||||
/// and CW = `β ⊕ M` (sign mux is a no-op under XOR).
|
||||
/// @details Splits the payload into leaf words (`naked(y0) ± naked(y1)`) so
|
||||
/// scalar `β = y0 + y1` (or `y0 ⊕ y1`) is never formed. The packing
|
||||
/// lane is muxed from XOR bit shares of the point, matching
|
||||
/// `mux_leaf_share` on the socket. Ring: `g = CCMP(t0,t1)` selects
|
||||
/// `N − M` vs `M − N` (matches `make_leaf` with `sign = t0`).
|
||||
/// Characteristic 2: CW = `N ⊕ M` (sign mux is a no-op under XOR).
|
||||
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
|
||||
/// @tparam I output index
|
||||
/// @tparam OutputsTuple outputs tuple
|
||||
/// @tparam InteriorBlock interior block
|
||||
/// @tparam OutputT output type
|
||||
/// @tparam InputT input domain type
|
||||
/// @param seed0 the `seed0`
|
||||
/// @param seed1 the `seed1`
|
||||
/// @param t0 the `t0`
|
||||
|
|
@ -557,13 +789,14 @@ struct local_cw_protocol
|
|||
/// @param y0 the `y0`
|
||||
/// @param y1 the `y1`
|
||||
/// @param pos_base the `pos_base`
|
||||
/// @param lane_x lane of the shared payload
|
||||
/// @return the opened leaf correction word
|
||||
/// @param x0 party 0's XOR share of the (lane) point
|
||||
/// @param x1 party 1's XOR share of the (lane) point
|
||||
/// @return the opened leaf correction word
|
||||
template <typename ExteriorPRG, std::size_t I = 0, typename OutputsTuple,
|
||||
typename InteriorBlock, typename OutputT>
|
||||
typename InteriorBlock, typename OutputT, typename InputT>
|
||||
auto open_arith_leaf(const InteriorBlock & seed0, const InteriorBlock & seed1,
|
||||
uint8_t t0, uint8_t t1, OutputT y0, OutputT y1, std::size_t pos_base,
|
||||
std::size_t lane_x) -> dpf::leaf_node_t<typename ExteriorPRG::block_type,
|
||||
InputT x0, InputT x1) -> dpf::leaf_node_t<typename ExteriorPRG::block_type,
|
||||
OutputT>
|
||||
{
|
||||
using output_type = OutputT;
|
||||
|
|
@ -573,36 +806,99 @@ struct local_cw_protocol
|
|||
using leaf_type = dpf::leaf_node_t<node_type, output_type>;
|
||||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||||
|
||||
constexpr std::size_t lg =
|
||||
dpf::lg_outputs_per_leaf_v<output_type, node_type>;
|
||||
constexpr auto to_int = utils::to_integral_type<InputT>{};
|
||||
auto bit0_at = [&](std::size_t b) {
|
||||
return static_cast<uint8_t>((to_int(x0) >> b) & 1u);
|
||||
};
|
||||
auto bit1_at = [&](std::size_t b) {
|
||||
return static_cast<uint8_t>((to_int(x1) >> b) & 1u);
|
||||
};
|
||||
auto naked_of = [&](output_type y) {
|
||||
return mux_naked_leaf_local<leaf_type>(lg,
|
||||
[&](unsigned i) {
|
||||
return dpf::make_naked_leaf<node_type>(
|
||||
static_cast<InputT>(i), y);
|
||||
},
|
||||
bit0_at, bit1_at);
|
||||
};
|
||||
// Split payload across leaf words; never form scalar β.
|
||||
const leaf_type naked =
|
||||
dpf::add_leaf<output_type>(naked_of(y0), naked_of(y1));
|
||||
|
||||
const auto M = dpf::make_leaf_mask<ExteriorPRG, I, OutputsTuple, InteriorBlock>(
|
||||
seed0, seed1, pos_base);
|
||||
output_type beta{};
|
||||
if constexpr (utils::has_characteristic_two_v<output_type>)
|
||||
{
|
||||
(void)t0;
|
||||
(void)t1;
|
||||
beta = static_cast<output_type>(y0 ^ y1);
|
||||
const leaf_type naked = dpf::make_naked_leaf<node_type>(lane_x, beta);
|
||||
return dpf::subtract_leaf<output_type>(naked, M);
|
||||
}
|
||||
else
|
||||
{
|
||||
const uint8_t g = open_ccmp(t0, t1);
|
||||
beta = static_cast<output_type>(y0 + y1);
|
||||
const leaf_type naked = dpf::make_naked_leaf<node_type>(lane_x, beta);
|
||||
// CW = (−1)^{t1}(β − M): g=0 → β−M; g=1 → M−β. Matches make_leaf(sign=t0).
|
||||
// CW = (−1)^{t1}(N − M): g=0 → N−M; g=1 → M−N. Matches make_leaf(sign=t0).
|
||||
if (g & 1u)
|
||||
return dpf::subtract_leaf<output_type>(M, naked);
|
||||
return dpf::subtract_leaf<output_type>(naked, M);
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Open a group of leaf correction words for one prefix group. In this
|
||||
/// local joint simulation both XOR shares of the point are present, so the
|
||||
/// point is reconstructed *inside* the protocol and handed to `leaf_fn`
|
||||
/// (which runs `make_leaves` for the group). The Doerner–Shelat gen never
|
||||
/// forms `x = x0 ^ x1` at its own call site; an MPC backend would instead
|
||||
/// run a per-group leaf CW exchange that never reveals `x`. After
|
||||
/// `encode_walk_shares`, arithmetic inputs are already `(alpha, 0)`.
|
||||
/// @brief Open the comparison threshold lane from XOR shares of the point.
|
||||
/// @details Reconstructs only inside this protocol hook for paint units,
|
||||
/// domain-edge triviality, and blocked suffixes. Per-level value
|
||||
/// words use share bits (`bit0 ^ bit1`) instead of this value.
|
||||
/// @tparam InputT input domain type
|
||||
/// @param x0 party 0's XOR share
|
||||
/// @param x1 party 1's XOR share
|
||||
/// @param nbits width of the comparison lane
|
||||
/// @return the comparison threshold as an integer lane
|
||||
template <typename InputT>
|
||||
unsigned __int128 open_cmp_threshold(InputT x0, InputT x1,
|
||||
std::size_t nbits) noexcept
|
||||
{
|
||||
constexpr auto to_int = utils::to_integral_type<InputT>{};
|
||||
constexpr std::size_t bl = utils::bitlength_of_v<InputT>;
|
||||
const InputT x = utils::xor_input_shares(x0, x1);
|
||||
if (nbits >= bl)
|
||||
return static_cast<unsigned __int128>(to_int(x));
|
||||
return static_cast<unsigned __int128>(to_int(x) >> (bl - nbits));
|
||||
}
|
||||
|
||||
/// @brief Public correction seed from XOR shares of the path prefix.
|
||||
/// @details Reconstructs the prefix only inside this protocol hook and
|
||||
/// returns `make_cs` (same digest as `oblivious_cs` when both
|
||||
/// seeds are in-process). The clear prefix is not returned.
|
||||
/// @tparam InputT input domain type
|
||||
/// @param level fold level (or blocked tag | level)
|
||||
/// @param x0 party 0's XOR share of the encoded point
|
||||
/// @param x1 party 1's XOR share of the encoded point
|
||||
/// @param bits number of high path bits in the prefix
|
||||
/// @param s0 party 0's on-path seed
|
||||
/// @param s1 party 1's on-path seed
|
||||
/// @return the public correction seed
|
||||
template <typename InputT>
|
||||
cs_block open_correction_seed(std::size_t level, InputT x0, InputT x1,
|
||||
std::size_t bits, simde__m128i s0, simde__m128i s1) noexcept
|
||||
{
|
||||
constexpr auto to_int = utils::to_integral_type<InputT>{};
|
||||
constexpr std::size_t bl = utils::bitlength_of_v<InputT>;
|
||||
const InputT x = utils::xor_input_shares(x0, x1);
|
||||
// Prefer `uint64_t` over `psnip_uint64_t{...}`: that macro expands to
|
||||
// `long unsigned int`, which is not a valid braced/cast type-id alone.
|
||||
const uint64_t prefix = (bits == 0 || bits > bl)
|
||||
? uint64_t{0}
|
||||
: static_cast<uint64_t>(to_int(x) >> (bl - bits));
|
||||
return detail::vdpf::make_cs(level, prefix, s0, s1);
|
||||
}
|
||||
|
||||
/// @brief Open a group of leaf correction words for one prefix group.
|
||||
/// @details Hands both XOR shares to `leaf_fn`. The builder may reconstruct
|
||||
/// the point only to emit public CWs (never return α to the DS
|
||||
/// call site). An MPC backend would run a per-group leaf CW
|
||||
/// exchange that never reveals `x`. Additive inputs have already
|
||||
/// been replaced by XOR shares of the sum.
|
||||
/// @tparam InputT input domain type
|
||||
/// @tparam LeafFn leaf fn
|
||||
/// @param x0 the `x0`
|
||||
|
|
@ -611,7 +907,7 @@ struct local_cw_protocol
|
|||
template <typename InputT, typename LeafFn>
|
||||
void open_leaf_group(InputT x0, InputT x1, LeafFn && leaf_fn)
|
||||
{
|
||||
std::forward<LeafFn>(leaf_fn)(utils::xor_input_shares(x0, x1));
|
||||
std::forward<LeafFn>(leaf_fn)(x0, x1);
|
||||
}
|
||||
};
|
||||
|
||||
|
|
@ -705,8 +1001,8 @@ void ds_advance_level(ds_gen_state<NodeT> & st, InputT x0, InputT x1,
|
|||
vblinds.R0 = v0[1];
|
||||
vblinds.L1 = v1[0];
|
||||
vblinds.R1 = v1[1];
|
||||
const int ai = static_cast<int>(
|
||||
(cmp->thresh >> (cmp->nbits - 1 - level)) & 1);
|
||||
// Path bit from share bits (threshold bit equals the walk bit).
|
||||
const int ai = static_cast<int>((bit0 ^ bit1) & 1u);
|
||||
if (cmp->paint)
|
||||
{
|
||||
const uint64_t unit = dcf_impl::paint_unit(cmp->kind, level,
|
||||
|
|
@ -890,30 +1186,31 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
const uint8_t t0 = static_cast<uint8_t>(sign0);
|
||||
const uint8_t t1 = static_cast<uint8_t>(dpf::get_lo_bit(parent1));
|
||||
|
||||
input_type x = utils::xor_input_shares(x0, x1);
|
||||
leaf_tuple leaves0{};
|
||||
leaf_tuple leaves1{};
|
||||
beaver_tuple beavers0{};
|
||||
beaver_tuple beavers1{};
|
||||
if (arith_out)
|
||||
{
|
||||
constexpr auto to_int = utils::to_integral_type<input_type>{};
|
||||
const std::size_t lane = static_cast<std::size_t>(to_int(x));
|
||||
auto cw = proto.template open_arith_leaf<ExteriorPRG, 0, outputs_tuple>(
|
||||
dpf::unset_lo_2bits(parent0), dpf::unset_lo_2bits(parent1), t0, t1,
|
||||
y0, y1, std::size_t{0}, lane);
|
||||
y0, y1, std::size_t{0}, x0, x1);
|
||||
std::get<0>(leaves0) = cw;
|
||||
std::get<0>(leaves1) = cw;
|
||||
}
|
||||
else
|
||||
{
|
||||
auto built = dpf::make_leaves<ExteriorPRG>(x,
|
||||
dpf::unset_lo_2bits(parent0), dpf::unset_lo_2bits(parent1), sign0,
|
||||
std::size_t{0}, y0);
|
||||
leaves0 = std::move(built.first.first);
|
||||
beavers0 = std::move(built.first.second);
|
||||
leaves1 = std::move(built.second.first);
|
||||
beavers1 = std::move(built.second.second);
|
||||
// Reconstruct the point only inside the leaf protocol hook.
|
||||
proto.open_leaf_group(x0, x1, [&](input_type sx0, input_type sx1) {
|
||||
const input_type x = utils::xor_input_shares(sx0, sx1);
|
||||
auto built = dpf::make_leaves<ExteriorPRG>(x,
|
||||
dpf::unset_lo_2bits(parent0), dpf::unset_lo_2bits(parent1),
|
||||
sign0, std::size_t{0}, y0);
|
||||
leaves0 = std::move(built.first.first);
|
||||
beavers0 = std::move(built.first.second);
|
||||
leaves1 = std::move(built.second.first);
|
||||
beavers1 = std::move(built.second.second);
|
||||
});
|
||||
(void)y1;
|
||||
}
|
||||
|
||||
|
|
@ -1000,18 +1297,29 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
const node parent1 = st.seed1();
|
||||
const bool sign0 = dpf::get_lo_bit(parent0);
|
||||
|
||||
input_type x = utils::xor_input_shares(x0, x1);
|
||||
auto built = dpf::make_leaves<ExteriorPRG>(x,
|
||||
dpf::unset_lo_2bits(parent0), dpf::unset_lo_2bits(parent1), sign0,
|
||||
std::size_t{0}, std::forward<OutputT>(y), std::forward<OutputTs>(ys)...);
|
||||
typename dpf_type::leaf_tuple leaves0{};
|
||||
typename dpf_type::beaver_tuple beavers0{};
|
||||
typename dpf_type::leaf_tuple leaves1{};
|
||||
typename dpf_type::beaver_tuple beavers1{};
|
||||
proto.open_leaf_group(x0, x1, [&](input_type sx0, input_type sx1) {
|
||||
const input_type x = utils::xor_input_shares(sx0, sx1);
|
||||
auto built = dpf::make_leaves<ExteriorPRG>(x,
|
||||
dpf::unset_lo_2bits(parent0), dpf::unset_lo_2bits(parent1), sign0,
|
||||
std::size_t{0}, std::forward<OutputT>(y),
|
||||
std::forward<OutputTs>(ys)...);
|
||||
leaves0 = std::move(built.first.first);
|
||||
beavers0 = std::move(built.first.second);
|
||||
leaves1 = std::move(built.second.first);
|
||||
beavers1 = std::move(built.second.second);
|
||||
});
|
||||
|
||||
input_type off0{};
|
||||
input_type off1{};
|
||||
return dpf::make_party_key_pair(
|
||||
dpf_type{root0, correction_words, correction_advice,
|
||||
built.first.first, built.first.second, off0},
|
||||
leaves0, beavers0, off0},
|
||||
dpf_type{root1, correction_words, correction_advice,
|
||||
built.second.first, built.second.second, off1});
|
||||
leaves1, beavers1, off1});
|
||||
}
|
||||
|
||||
/// @brief Single-output plaintext β (disambiguates from arith_out overload).
|
||||
|
|
|
|||
775
include/dpf/dpf3.hpp
Normal file
775
include/dpf/dpf3.hpp
Normal file
|
|
@ -0,0 +1,775 @@
|
|||
/// @file dpf/dpf3.hpp
|
||||
/// @brief Three-evaluator (2,3) point DPF after ePrint 2024/1658.
|
||||
/// @details Each party key is a pair of two-party VDPF+ keys. Evaluation is
|
||||
/// two ordinary walks, an XOR, and a party-index scale in `fp61`.
|
||||
/// Reconstruction is Shamir interpolation.
|
||||
/// @note Following Zyskind, Yanai, and Pentland, ePrint 2024/1658, Figure 3: each evaluator holds one key from each of two (2,2)-VDPF+ instances. Their evaluation section records about 2× the key size of one two-party DPF.
|
||||
///
|
||||
/// **Updatable keys** (`dpf::updatable`) keep beaver-backed XOR leaves
|
||||
/// so `update_payload` can rewrite `β` with four leaf patches and a
|
||||
/// refresh of the public offsets `π` — `O(λ)`, independent of the
|
||||
/// domain — without moving `α`. Non-updatable keys bake the leaf;
|
||||
/// calling `update_payload` on them throws.
|
||||
/// @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_DPF3_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_DPF3_HPP__
|
||||
|
||||
#include <array>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <stdexcept>
|
||||
#include <tuple>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/eval_full.hpp"
|
||||
#include "dpf/eval_interval.hpp"
|
||||
#include "dpf/eval_point.hpp"
|
||||
#include "dpf/eval_sequence.hpp"
|
||||
#include "dpf/fp61.hpp"
|
||||
#include "dpf/leaf_node.hpp"
|
||||
#include "dpf/path_memoizer.hpp"
|
||||
#include "dpf/placement.hpp"
|
||||
#include "dpf/prg_aes.hpp"
|
||||
#include "dpf/random.hpp"
|
||||
#include "dpf/shamir3.hpp"
|
||||
#include "dpf/utils.hpp"
|
||||
#include "dpf/verifiable.hpp"
|
||||
#include "dpf/wildcard.hpp"
|
||||
#include "dpf/xor_wrapper.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
/// @brief Phantom tag selecting the three-evaluator point construction.
|
||||
struct dpf3_t
|
||||
{
|
||||
static constexpr bool is_dpf3_tag = true;
|
||||
};
|
||||
|
||||
inline constexpr dpf3_t dpf3{};
|
||||
|
||||
namespace detail
|
||||
{
|
||||
namespace dpf3_impl
|
||||
{
|
||||
|
||||
using xor61 = shamir3::xor61;
|
||||
|
||||
template <typename Inner>
|
||||
struct vdpf_plus_key
|
||||
{
|
||||
using inner_type = Inner;
|
||||
using input_type = typename Inner::input_type;
|
||||
/// @brief Inner two-party spine. Eval that accepts a `dpf_key` also accepts
|
||||
/// this object and reads `dpf_key`.
|
||||
Inner dpf_key{};
|
||||
xor61 offset{};
|
||||
};
|
||||
|
||||
template <typename Inner, typename PathMemoizer = basic_path_memoizer<Inner>>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
xor61 eval_plus(const vdpf_plus_key<Inner> & key, typename Inner::input_type x,
|
||||
PathMemoizer && path = PathMemoizer{})
|
||||
{
|
||||
const auto y = *dpf::eval_point(key.dpf_key, x,
|
||||
std::forward<PathMemoizer>(path));
|
||||
if constexpr (is_secret_share_v<std::decay_t<decltype(y)>>)
|
||||
return xor61{y.raw()} + key.offset;
|
||||
else
|
||||
return xor61{static_cast<std::uint64_t>(y)} + key.offset;
|
||||
}
|
||||
|
||||
template <typename Inner, typename PathMemoizer = basic_path_memoizer<Inner>>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
xor61 eval_plus(const vdpf_plus_key<Inner> & key, typename Inner::input_type x,
|
||||
prove_ref pr, PathMemoizer && path = PathMemoizer{})
|
||||
{
|
||||
const auto y = *dpf::eval_point(key.dpf_key, x, pr,
|
||||
std::forward<PathMemoizer>(path));
|
||||
// Re-bind the public offset (refreshed by `update_payload`).
|
||||
const auto off = static_cast<std::uint64_t>(key.offset);
|
||||
detail::vdpf::fold_bytes(pr.token, 0x50, &off, sizeof(off));
|
||||
if constexpr (is_secret_share_v<std::decay_t<decltype(y)>>)
|
||||
return xor61{y.raw()} + key.offset;
|
||||
else
|
||||
return xor61{static_cast<std::uint64_t>(y)} + key.offset;
|
||||
}
|
||||
|
||||
template <typename Buf>
|
||||
xor61 xor61_from_buf_elem(const Buf & e)
|
||||
{
|
||||
if constexpr (is_secret_share_v<std::decay_t<Buf>>)
|
||||
return xor61{e.raw()};
|
||||
else
|
||||
return xor61{static_cast<std::uint64_t>(e)};
|
||||
}
|
||||
|
||||
template <typename KeyT, typename BufA, typename BufB>
|
||||
void combine_spine_bufs(const KeyT & key, const BufA & a, const BufB & b,
|
||||
std::vector<fp61> & out)
|
||||
{
|
||||
const std::size_t n = a.size();
|
||||
out.resize(n);
|
||||
const fp61 scale{static_cast<std::uint64_t>(KeyT::party)};
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
const xor61 y = xor61_from_buf_elem(a[i]) + key.a.offset
|
||||
+ xor61_from_buf_elem(b[i]) + key.b.offset;
|
||||
out[i] = shamir3::xor_scale(y, scale);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename Inner>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
xor61 peel(const Inner & key, typename Inner::input_type x)
|
||||
{
|
||||
const auto y = *dpf::eval_point(key, x);
|
||||
if constexpr (is_secret_share_v<std::decay_t<decltype(y)>>)
|
||||
return xor61{y.raw()};
|
||||
else
|
||||
return xor61{static_cast<std::uint64_t>(y)};
|
||||
}
|
||||
|
||||
struct tau_quad
|
||||
{
|
||||
xor61 t0{};
|
||||
xor61 t1{};
|
||||
xor61 t2{};
|
||||
xor61 t3{};
|
||||
};
|
||||
|
||||
inline tau_quad sample_taus(fp61 beta)
|
||||
{
|
||||
const auto shares = shamir3::share_secret(beta);
|
||||
const fp61 s1 = shamir3::unscale(shares[0]);
|
||||
const fp61 s2 = shamir3::unscale(shares[1]);
|
||||
const fp61 s3 = shamir3::unscale(shares[2]);
|
||||
tau_quad t{};
|
||||
for (int attempt = 0; attempt < 16; ++attempt)
|
||||
{
|
||||
t.t0 = xor61{uniform_sample<std::uint64_t>() & fp61_mod};
|
||||
t.t2 = shamir3::field_xor(s1, t.t0);
|
||||
t.t1 = shamir3::field_xor(s2, t.t2);
|
||||
t.t3 = shamir3::field_xor(s3, t.t1);
|
||||
const auto ok = [](xor61 w) {
|
||||
return (static_cast<std::uint64_t>(w) & fp61_mod) != fp61_mod;
|
||||
};
|
||||
if (ok(t.t0) && ok(t.t1) && ok(t.t2) && ok(t.t3))
|
||||
return t;
|
||||
}
|
||||
throw std::runtime_error("make_dpf3: embed resampling failed");
|
||||
}
|
||||
|
||||
template <typename Key, typename Output = xor61>
|
||||
void patch_leaf_xor(Key & key, typename Key::input_type alpha, Output delta)
|
||||
{
|
||||
using node = typename Key::exterior_node;
|
||||
using concrete = dpf::concrete_type_t<Output>;
|
||||
auto & wrap = std::get<0>(key.leaf_nodes);
|
||||
auto & leaf = wrap.raw_leaf();
|
||||
leaf = dpf::add_leaf<concrete>(leaf,
|
||||
dpf::make_naked_leaf<node>(alpha, concrete{delta}));
|
||||
}
|
||||
|
||||
template <typename K0, typename K1, typename Share0, typename Share1>
|
||||
void assign_wildcard_pair(K0 & k0, K1 & k1, Share0 share0, Share1 share1)
|
||||
{
|
||||
auto & w0 = std::get<0>(k0.leaf_nodes);
|
||||
auto & w1 = std::get<0>(k1.leaf_nodes);
|
||||
if (w0.is_ready())
|
||||
w0.begin_update();
|
||||
if (w1.is_ready())
|
||||
w1.begin_update();
|
||||
const auto b0 = w0.compute_and_get_blinded_output_share(share0);
|
||||
const auto b1 = w1.compute_and_get_blinded_output_share(share1);
|
||||
const auto l0 = w0.compute_and_get_leaf_share(b1);
|
||||
const auto l1 = w1.compute_and_get_leaf_share(b0);
|
||||
w0.reconstruct_correction_word(l1);
|
||||
w1.reconstruct_correction_word(l0);
|
||||
}
|
||||
|
||||
template <typename K0, typename K1>
|
||||
void assign_xor_payload(K0 & k0, K1 & k1, xor61 payload)
|
||||
{
|
||||
const xor61 s0{uniform_sample<std::uint64_t>()};
|
||||
const xor61 s1 = payload + s0;
|
||||
assign_wildcard_pair(k0, k1, s0, s1);
|
||||
}
|
||||
|
||||
} // namespace dpf3_impl
|
||||
} // namespace detail
|
||||
|
||||
/// @brief One evaluator's key in a (2,3) point DPF.
|
||||
/// @tparam Party party index in `{1, 2, 3}`
|
||||
/// @tparam PlusA VDPF+ key type for instance A
|
||||
/// @tparam PlusB VDPF+ key type for instance B
|
||||
template <int Party, typename PlusA, typename PlusB>
|
||||
struct dpf3_key
|
||||
{
|
||||
static_assert(Party >= 1 && Party <= 3, "dpf3 party is 1, 2, or 3");
|
||||
static constexpr int party = Party;
|
||||
static constexpr bool is_dpf3 = true;
|
||||
using input_type = typename PlusA::input_type;
|
||||
using plus_a_type = PlusA;
|
||||
using plus_b_type = PlusB;
|
||||
|
||||
PlusA a{};
|
||||
PlusB b{};
|
||||
bool verifiable = false;
|
||||
bool extractable = false;
|
||||
bool updatable = false;
|
||||
};
|
||||
|
||||
namespace detail
|
||||
{
|
||||
namespace dpf3_impl
|
||||
{
|
||||
|
||||
/// @brief Open the XOR-shared point the same way local DS walks it.
|
||||
template <typename InputT>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
InputT open_xor_point(InputT x0, InputT x1)
|
||||
{
|
||||
utils::flip_msb_if_signed_integral(x0);
|
||||
constexpr auto to_int = utils::to_integral_type<InputT>{};
|
||||
using I = decltype(to_int(x0));
|
||||
return utils::make_from_integral_value<InputT>{}(
|
||||
static_cast<I>(to_int(x0) ^ to_int(x1)));
|
||||
}
|
||||
|
||||
/// @brief Pack Fig-3 party keys from two completed two-party spines + `τ`.
|
||||
/// @details Computes public `π` from peels of the party-0 halves at `α`.
|
||||
/// Parameter names avoid `B0`/`B1` (termios baud macros).
|
||||
template <bool Verifiable, bool Extractable, bool Updatable, typename KeyA0,
|
||||
typename KeyA1, typename KeyB0, typename KeyB1, typename Input>
|
||||
auto assemble_from_spines(KeyA0 key_a0, KeyA1 key_a1, KeyB0 key_b0,
|
||||
KeyB1 key_b1, tau_quad t, Input alpha)
|
||||
{
|
||||
using X = xor61;
|
||||
const X yA0 = peel(key_a0, alpha);
|
||||
const X yB0 = peel(key_b0, alpha);
|
||||
const X piA = t.t0 + yA0;
|
||||
const X piB = t.t2 + yB0;
|
||||
using PlusA0 = vdpf_plus_key<KeyA0>;
|
||||
using PlusA1 = vdpf_plus_key<KeyA1>;
|
||||
using PlusB0 = vdpf_plus_key<KeyB0>;
|
||||
using PlusB1 = vdpf_plus_key<KeyB1>;
|
||||
PlusA0 plus_a0{std::move(key_a0), piA};
|
||||
PlusA1 plus_a1{std::move(key_a1), piA};
|
||||
PlusB0 plus_b0{std::move(key_b0), piB};
|
||||
PlusB1 plus_b1{std::move(key_b1), piB};
|
||||
dpf3_key<1, PlusA0, PlusB0> k1{plus_a0, plus_b0, Verifiable, Extractable,
|
||||
Updatable};
|
||||
dpf3_key<2, PlusA1, PlusB0> k2{plus_a1, plus_b0, Verifiable, Extractable,
|
||||
Updatable};
|
||||
dpf3_key<3, PlusA1, PlusB1> k3{plus_a1, plus_b1, Verifiable, Extractable,
|
||||
Updatable};
|
||||
return std::make_tuple(std::move(k1), std::move(k2), std::move(k3));
|
||||
}
|
||||
|
||||
template <typename Input, typename InteriorPRG, typename ExteriorPRG,
|
||||
bool Verifiable, bool Extractable, bool Updatable>
|
||||
auto make_point3(Input alpha, fp61 beta)
|
||||
{
|
||||
using X = xor61;
|
||||
const tau_quad t = sample_taus(beta);
|
||||
const X payload_a = t.t0 + t.t1;
|
||||
const X payload_b = t.t2 + t.t3;
|
||||
|
||||
if constexpr (Updatable)
|
||||
{
|
||||
auto A = [&] {
|
||||
if constexpr (Verifiable)
|
||||
return dpf::make_dpf<InteriorPRG, ExteriorPRG>(alpha,
|
||||
dpf::wildcard_value<X>{}, dpf::verifiable{});
|
||||
else
|
||||
return dpf::make_dpf<InteriorPRG, ExteriorPRG>(alpha,
|
||||
dpf::wildcard_value<X>{});
|
||||
}();
|
||||
auto B = [&] {
|
||||
if constexpr (Verifiable)
|
||||
return dpf::make_dpf<InteriorPRG, ExteriorPRG>(alpha,
|
||||
dpf::wildcard_value<X>{}, dpf::verifiable{});
|
||||
else
|
||||
return dpf::make_dpf<InteriorPRG, ExteriorPRG>(alpha,
|
||||
dpf::wildcard_value<X>{});
|
||||
}();
|
||||
assign_xor_payload(A.first, A.second, payload_a);
|
||||
assign_xor_payload(B.first, B.second, payload_b);
|
||||
return assemble_from_spines<Verifiable, Extractable, Updatable>(
|
||||
std::move(A.first), std::move(A.second), std::move(B.first),
|
||||
std::move(B.second), t, alpha);
|
||||
}
|
||||
else
|
||||
{
|
||||
auto A = [&] {
|
||||
if constexpr (Verifiable)
|
||||
return dpf::make_dpf<InteriorPRG, ExteriorPRG>(alpha, payload_a,
|
||||
dpf::verifiable{});
|
||||
else
|
||||
return dpf::make_dpf<InteriorPRG, ExteriorPRG>(alpha, payload_a);
|
||||
}();
|
||||
auto B = [&] {
|
||||
if constexpr (Verifiable)
|
||||
return dpf::make_dpf<InteriorPRG, ExteriorPRG>(alpha, payload_b,
|
||||
dpf::verifiable{});
|
||||
else
|
||||
return dpf::make_dpf<InteriorPRG, ExteriorPRG>(alpha, payload_b);
|
||||
}();
|
||||
return assemble_from_spines<Verifiable, Extractable, Updatable>(
|
||||
std::move(A.first), std::move(A.second), std::move(B.first),
|
||||
std::move(B.second), t, alpha);
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Flags carried by `verifiable` / `extractable` / `updatable` tags.
|
||||
/// @details Any subset, any order. A repeated tag is rejected.
|
||||
template <typename ...Tags>
|
||||
struct tag_flags
|
||||
{
|
||||
static constexpr bool verifiable =
|
||||
(is_verifiable_tag_v<std::decay_t<Tags>> || ...);
|
||||
static constexpr bool extractable =
|
||||
(is_extractable_tag_v<std::decay_t<Tags>> || ...);
|
||||
static constexpr bool updatable =
|
||||
(is_updatable_tag_v<std::decay_t<Tags>> || ...);
|
||||
static constexpr bool known = ((is_verifiable_tag_v<std::decay_t<Tags>>
|
||||
|| is_extractable_tag_v<std::decay_t<Tags>>
|
||||
|| is_updatable_tag_v<std::decay_t<Tags>>) && ...);
|
||||
static constexpr std::size_t counted =
|
||||
static_cast<std::size_t>(verifiable)
|
||||
+ static_cast<std::size_t>(extractable)
|
||||
+ static_cast<std::size_t>(updatable);
|
||||
};
|
||||
|
||||
template <bool Verifiable, bool Extractable, bool Updatable,
|
||||
typename InteriorPRG, typename ExteriorPRG, typename Input>
|
||||
auto make_tagged(Input alpha, fp61 beta)
|
||||
{
|
||||
return make_point3<Input, InteriorPRG, ExteriorPRG, Verifiable, Extractable,
|
||||
Updatable>(alpha, beta);
|
||||
}
|
||||
|
||||
/// @brief Read the four planted τ strings from live VDPF+ evaluations at `α`.
|
||||
template <typename K1, typename K2, typename K3, typename Input>
|
||||
tau_quad read_taus(const K1 & k1, const K2 & k2, const K3 & k3, Input alpha)
|
||||
{
|
||||
tau_quad t{};
|
||||
t.t0 = peel(k1.a.dpf_key, alpha) + k1.a.offset;
|
||||
t.t1 = peel(k2.a.dpf_key, alpha) + k2.a.offset;
|
||||
t.t2 = peel(k1.b.dpf_key, alpha) + k1.b.offset;
|
||||
t.t3 = peel(k3.b.dpf_key, alpha) + k3.b.offset;
|
||||
return t;
|
||||
}
|
||||
|
||||
template <typename Plus, typename Input>
|
||||
void refresh_offset(Plus & plus, Input alpha, xor61 target_delta0)
|
||||
{
|
||||
plus.offset = target_delta0 + peel(plus.dpf_key, alpha);
|
||||
}
|
||||
|
||||
/// @brief XOR a precomputed naked-leaf patch onto a ready inner key.
|
||||
template <typename Key, typename Leaf>
|
||||
void apply_leaf_patch(Key & key, const Leaf & patch)
|
||||
{
|
||||
using concrete = dpf::concrete_type_t<typename Key::template output_type_t<0>>;
|
||||
auto & wrap = std::get<0>(key.leaf_nodes);
|
||||
auto & leaf = wrap.raw_leaf();
|
||||
leaf = dpf::add_leaf<concrete>(leaf, patch);
|
||||
}
|
||||
|
||||
/// @brief Build the naked-leaf Fig-10 patch for payload difference `delta`.
|
||||
template <typename Key, typename Input>
|
||||
auto make_leaf_patch(Input alpha, xor61 delta)
|
||||
{
|
||||
using node = typename Key::exterior_node;
|
||||
using concrete = dpf::concrete_type_t<typename Key::template output_type_t<0>>;
|
||||
return dpf::make_naked_leaf<node>(alpha, concrete{delta});
|
||||
}
|
||||
|
||||
} // namespace dpf3_impl
|
||||
} // namespace detail
|
||||
|
||||
/// @brief Generate three (2,3) point keys for `f(α) = β`.
|
||||
/// @details Optional tags are `verifiable`, `extractable`, and `updatable`,
|
||||
/// in any order. `extractable` is the outer fp61 sketch flag; inner
|
||||
/// XOR keys stay ordinary. `updatable` keeps beaver leaves so
|
||||
/// `update_payload` can rewrite `β`.
|
||||
/// @note Following Zyskind, Yanai, and Pentland, ePrint 2024/1658, Figure 3: two (2,2)-VDPF+ spines and two walks.
|
||||
/// \complexity O(n) time. Two `make_dpf` spines (A and B), each the point-keygen loop, plus a constant number of `eval_point` peels in `assemble_from_spines`. No messages.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
typename ...Tags>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto make_dpf3(InputT alpha, fp61 beta, Tags ...tags)
|
||||
{
|
||||
using flags = detail::dpf3_impl::tag_flags<Tags...>;
|
||||
static_assert(flags::known, "make_dpf3 tags are verifiable, extractable, updatable");
|
||||
static_assert(sizeof...(Tags) == flags::counted,
|
||||
"make_dpf3: repeated tag");
|
||||
(void)std::initializer_list<int>{((void)tags, 0)...};
|
||||
return detail::dpf3_impl::make_tagged<flags::verifiable, flags::extractable,
|
||||
flags::updatable, InteriorPRG, ExteriorPRG>(alpha, beta);
|
||||
}
|
||||
|
||||
/// @brief Evaluate one party's (2,3) key at `x`.
|
||||
/// \complexity O(n) time. n is `depth`. One interior traversal per level from the memoizer resume index through the leaf. Extra space is the path memoizer (O(n) nodes, or one node if it does not memoize). A proof token adds one fold per level walked.
|
||||
template <typename KeyT, typename Query,
|
||||
typename PathA = basic_path_memoizer<
|
||||
typename KeyT::plus_a_type::inner_type>,
|
||||
typename PathB = basic_path_memoizer<
|
||||
typename KeyT::plus_b_type::inner_type>,
|
||||
std::enable_if_t<std::decay_t<KeyT>::is_dpf3, int> = 0>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
fp61 eval_point(const KeyT & key, Query && x, PathA && path_a = PathA{},
|
||||
PathB && path_b = PathB{})
|
||||
{
|
||||
const auto qx = static_cast<typename KeyT::input_type>(x);
|
||||
const auto ya = detail::dpf3_impl::eval_plus(key.a, qx,
|
||||
std::forward<PathA>(path_a));
|
||||
const auto yb = detail::dpf3_impl::eval_plus(key.b, qx,
|
||||
std::forward<PathB>(path_b));
|
||||
return shamir3::xor_scale(ya + yb,
|
||||
fp61{static_cast<std::uint64_t>(KeyT::party)});
|
||||
}
|
||||
|
||||
/// @brief Evaluate and fold an inner proof token from each VDPF+.
|
||||
/// \complexity O(n) time. n is `depth`. One interior traversal per level from the memoizer resume index through the leaf. Extra space is the path memoizer (O(n) nodes, or one node if it does not memoize). A proof token adds one fold per level walked.
|
||||
template <typename KeyT, typename Query,
|
||||
typename PathA = basic_path_memoizer<
|
||||
typename KeyT::plus_a_type::inner_type>,
|
||||
typename PathB = basic_path_memoizer<
|
||||
typename KeyT::plus_b_type::inner_type>,
|
||||
std::enable_if_t<std::decay_t<KeyT>::is_dpf3, int> = 0>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
fp61 eval_point(const KeyT & key, Query && x, prove_ref pr,
|
||||
PathA && path_a = PathA{}, PathB && path_b = PathB{})
|
||||
{
|
||||
if (!key.verifiable)
|
||||
throw std::invalid_argument("eval_point(prove): key is not verifiable");
|
||||
proof_token pa{}, pb{};
|
||||
const auto qx = static_cast<typename KeyT::input_type>(x);
|
||||
const auto ya = detail::dpf3_impl::eval_plus(key.a, qx, prove(pa),
|
||||
std::forward<PathA>(path_a));
|
||||
const auto yb = detail::dpf3_impl::eval_plus(key.b, qx, prove(pb),
|
||||
std::forward<PathB>(path_b));
|
||||
pr.token = detail::vdpf::xor_proof(pa, pb);
|
||||
return shamir3::xor_scale(ya + yb,
|
||||
fp61{static_cast<std::uint64_t>(KeyT::party)});
|
||||
}
|
||||
|
||||
/// @brief Full-domain (2,3) eval: expand each spine once, then XOR and scale.
|
||||
/// \complexity Same expansion as `eval_interval` on the whole domain. L = 2^{n - lg(outputs_per_leaf)} leaf nodes, n = `depth`. Time Θ(L) interior traversals. The output buffer stores one slot per domain point (2^n).
|
||||
template <typename KeyT,
|
||||
std::enable_if_t<std::decay_t<KeyT>::is_dpf3, int> = 0>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::vector<fp61> eval_full(const KeyT & key)
|
||||
{
|
||||
auto buf_a = dpf::make_output_buffer_for_full(key.a.dpf_key);
|
||||
auto buf_b = dpf::make_output_buffer_for_full(key.b.dpf_key);
|
||||
dpf::eval_full(key.a.dpf_key, buf_a);
|
||||
dpf::eval_full(key.b.dpf_key, buf_b);
|
||||
std::vector<fp61> out;
|
||||
detail::dpf3_impl::combine_spine_bufs(key, buf_a, buf_b, out);
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief Add a (2,3) full-domain expansion into a caller's share vector.
|
||||
/// @details `buf[i] += eval_full(key)[i]` for every slot. Shamir shares are
|
||||
/// linear, so summing appends per party and reconstructing any two
|
||||
/// recovers the running total. A (2,3) ledger folds each append with
|
||||
/// this call instead of an `eval_point` loop.
|
||||
/// \complexity Same expansion as `eval_full` on the (2,3) key.
|
||||
template <typename KeyT, typename Buffer,
|
||||
std::enable_if_t<std::decay_t<KeyT>::is_dpf3, int> = 0>
|
||||
void eval_full_add_into(Buffer & buf, const KeyT & key) // NOLINT(runtime/references)
|
||||
{
|
||||
const auto full = eval_full(key);
|
||||
const std::size_t n = std::min<std::size_t>(full.size(), buf.size());
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
buf[i] = buf[i] + full[i];
|
||||
}
|
||||
|
||||
/// @brief Dot a (2,3) full-domain expansion with a public table.
|
||||
/// @details `sum_i eval_full(key)[i] * weights[i]`. Shamir shares are linear,
|
||||
/// so any two parties' dots reconstruct the table entry at `α` when
|
||||
/// the payload is `1`. A three-server PIR is this call per server.
|
||||
/// \complexity Same expansion as `eval_full` on the (2,3) key, plus one
|
||||
/// multiply-add per domain point.
|
||||
template <typename KeyT, typename Weights,
|
||||
std::enable_if_t<std::decay_t<KeyT>::is_dpf3, int> = 0>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
fp61 eval_full_inner_product(const KeyT & key, const Weights & weights)
|
||||
{
|
||||
const auto full = eval_full(key);
|
||||
fp61 acc{};
|
||||
const std::size_t n = std::min<std::size_t>(full.size(), weights.size());
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
const auto & w = weights[i];
|
||||
if constexpr (std::is_same_v<std::decay_t<decltype(w)>, fp61>)
|
||||
acc = acc + full[i] * w;
|
||||
else
|
||||
acc = acc + full[i] * fp61{static_cast<std::uint64_t>(w)};
|
||||
}
|
||||
return acc;
|
||||
}
|
||||
|
||||
/// @brief Interval (2,3) eval into an `fp61` buffer.
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <typename KeyT, typename LaneT,
|
||||
std::enable_if_t<std::decay_t<KeyT>::is_dpf3, int> = 0>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::vector<fp61> eval_interval(const KeyT & key, LaneT from, LaneT to)
|
||||
{
|
||||
auto buf_a = dpf::make_output_buffer(key.a.dpf_key, from, to);
|
||||
auto buf_b = dpf::make_output_buffer(key.b.dpf_key, from, to);
|
||||
dpf::eval_interval(key.a.dpf_key, from, to, buf_a);
|
||||
dpf::eval_interval(key.b.dpf_key, from, to, buf_b);
|
||||
std::vector<fp61> out;
|
||||
detail::dpf3_impl::combine_spine_bufs(key, buf_a, buf_b, out);
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief Pack an evaluation as a typed Shamir share.
|
||||
template <typename KeyT, std::enable_if_t<KeyT::is_dpf3, int> = 0>
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_CONST
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr shamir3::share as_share(const KeyT &, fp61 y) noexcept
|
||||
{
|
||||
return shamir3::share{KeyT::party, y};
|
||||
}
|
||||
|
||||
/// @brief The same evaluation as a party-tagged (2,3) Shamir share.
|
||||
/// @details `dpf3` parties are `1`, `2`, `3`. The typed share's party is one less.
|
||||
template <typename KeyT, std::enable_if_t<KeyT::is_dpf3, int> = 0>
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_CONST
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr shamir_share<fp61, static_cast<std::size_t>(KeyT::party - 1)>
|
||||
as_shamir_share(const KeyT &, fp61 y) noexcept
|
||||
{
|
||||
return shamir_share<fp61, static_cast<std::size_t>(KeyT::party - 1)>::from_raw(y);
|
||||
}
|
||||
|
||||
/// @brief Reconstruct from any two (2,3) evaluation shares.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline fp61 reconstruct(shamir3::share a, shamir3::share b)
|
||||
{
|
||||
return shamir3::reconstruct(a, b);
|
||||
}
|
||||
|
||||
/// @brief Reconstruct from all three shares.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline fp61 reconstruct(shamir3::share a, shamir3::share b, shamir3::share c)
|
||||
{
|
||||
return shamir3::reconstruct(a, b, c);
|
||||
}
|
||||
|
||||
/// @brief Three-party proof token: two inner tokens plus the public offsets.
|
||||
struct dpf3_proof
|
||||
{
|
||||
proof_token a{};
|
||||
proof_token b{};
|
||||
shamir3::xor61 offset_a{};
|
||||
shamir3::xor61 offset_b{};
|
||||
};
|
||||
|
||||
/// @brief Build a three-party proof at `x`.
|
||||
template <typename KeyT, typename Query,
|
||||
std::enable_if_t<KeyT::is_dpf3, int> = 0>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
dpf3_proof prove_dpf3(const KeyT & key, Query && x)
|
||||
{
|
||||
if (!key.verifiable)
|
||||
throw std::invalid_argument("prove_dpf3: key is not verifiable");
|
||||
dpf3_proof out{};
|
||||
out.offset_a = key.a.offset;
|
||||
out.offset_b = key.b.offset;
|
||||
const auto qx = static_cast<typename KeyT::input_type>(x);
|
||||
std::ignore = detail::dpf3_impl::eval_plus(key.a, qx, prove(out.a));
|
||||
std::ignore = detail::dpf3_impl::eval_plus(key.b, qx, prove(out.b));
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief Verify three (2,3) proofs agree on offsets and inner tokens.
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_PURE
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
bool verify_dpf3(const dpf3_proof & p1, const dpf3_proof & p2,
|
||||
const dpf3_proof & p3) noexcept
|
||||
{
|
||||
if (p1.offset_a != p2.offset_a || p2.offset_a != p3.offset_a)
|
||||
return false;
|
||||
if (p1.offset_b != p2.offset_b || p2.offset_b != p3.offset_b)
|
||||
return false;
|
||||
if (!verify(p1.a, p2.a))
|
||||
return false;
|
||||
if (!verify(p2.a, p3.a))
|
||||
return false;
|
||||
if (!verify(p1.b, p2.b))
|
||||
return false;
|
||||
if (!verify(p2.b, p3.b))
|
||||
return false;
|
||||
return true;
|
||||
}
|
||||
|
||||
/// @brief In-place payload update of three updatable (2,3) keys (Fig. 10).
|
||||
/// @details Reads the live `τ` strings from evaluations at `α`, samples a
|
||||
/// fresh Shamir split of `β'`, patches the four inner XOR leaves by
|
||||
/// the payload difference, refreshes `π`, and leaves the tree path
|
||||
/// untouched. Requires keys generated with `dpf::updatable`.
|
||||
/// @tparam K1 party-1 key type
|
||||
/// @tparam K2 party-2 key type
|
||||
/// @tparam K3 party-3 key type
|
||||
/// @tparam InputT input domain type
|
||||
/// @param k1 party 1 key
|
||||
/// @param k2 party 2 key
|
||||
/// @param k3 party 3 key
|
||||
/// @param alpha the same secret point
|
||||
/// @param beta_new the new payload
|
||||
/// @throws std::invalid_argument if any key is not updatable
|
||||
template <typename K1, typename K2, typename K3, typename InputT>
|
||||
void update_payload(K1 & k1, K2 & k2, K3 & k3, InputT alpha, fp61 beta_new)
|
||||
{
|
||||
static_assert(K1::is_dpf3 && K2::is_dpf3 && K3::is_dpf3, "dpf3 keys");
|
||||
if (!k1.updatable || !k2.updatable || !k3.updatable)
|
||||
throw std::invalid_argument(
|
||||
"update_payload: keys were not generated with dpf::updatable");
|
||||
|
||||
using X = detail::dpf3_impl::xor61;
|
||||
const auto told = detail::dpf3_impl::read_taus(k1, k2, k3, alpha);
|
||||
const auto tnew = detail::dpf3_impl::sample_taus(beta_new);
|
||||
const X dA = (tnew.t0 + tnew.t1) + (told.t0 + told.t1);
|
||||
const X dB = (tnew.t2 + tnew.t3) + (told.t2 + told.t3);
|
||||
|
||||
// After Beaver assign both parties hold the same leaf CW. Patch every
|
||||
// copy of each spine's CW by the payload difference (Fig. 10).
|
||||
// A0 on p1; A1 on p2 and p3.
|
||||
detail::dpf3_impl::patch_leaf_xor(k1.a.dpf_key, alpha, dA);
|
||||
detail::dpf3_impl::patch_leaf_xor(k2.a.dpf_key, alpha, dA);
|
||||
detail::dpf3_impl::patch_leaf_xor(k3.a.dpf_key, alpha, dA);
|
||||
// B0 on p1 and p2; B1 on p3.
|
||||
detail::dpf3_impl::patch_leaf_xor(k1.b.dpf_key, alpha, dB);
|
||||
detail::dpf3_impl::patch_leaf_xor(k2.b.dpf_key, alpha, dB);
|
||||
detail::dpf3_impl::patch_leaf_xor(k3.b.dpf_key, alpha, dB);
|
||||
|
||||
// π is public and identical on both halves of each VDPF+.
|
||||
// Leaf patches and the refreshed offset are re-bound on the next prove
|
||||
// (`init_proof` folds the leaf CW; `eval_plus` folds the offset).
|
||||
detail::dpf3_impl::refresh_offset(k1.a, alpha, tnew.t0);
|
||||
k2.a.offset = k1.a.offset;
|
||||
k3.a.offset = k1.a.offset;
|
||||
detail::dpf3_impl::refresh_offset(k1.b, alpha, tnew.t2);
|
||||
k2.b.offset = k1.b.offset;
|
||||
k3.b.offset = k1.b.offset;
|
||||
}
|
||||
|
||||
/// @brief Weight-1 sketch over three Shamir full-domain vectors (ungated).
|
||||
/// @details Prefer the key-taking overload, which enforces `dpf::extractable`.
|
||||
template <typename RRange>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
bool sketch_verify3(const std::vector<fp61> & s1, const std::vector<fp61> & s2,
|
||||
const std::vector<fp61> & s3, RRange && challenges)
|
||||
{
|
||||
if (s1.size() != s2.size() || s2.size() != s3.size())
|
||||
return false;
|
||||
std::vector<fp61> opened;
|
||||
opened.reserve(s1.size());
|
||||
std::vector<fp61> rs;
|
||||
rs.reserve(s1.size());
|
||||
std::size_t i = 0;
|
||||
for (auto && r : challenges)
|
||||
{
|
||||
if (i >= s1.size())
|
||||
return false;
|
||||
opened.push_back(shamir3::reconstruct(
|
||||
shamir3::share{1, s1[i]}, shamir3::share{2, s2[i]},
|
||||
shamir3::share{3, s3[i]}));
|
||||
rs.push_back(r);
|
||||
++i;
|
||||
}
|
||||
if (i != s1.size())
|
||||
return false;
|
||||
sketch_share sk = sketch_fold(opened, rs);
|
||||
sketch_share zero{};
|
||||
return sketch_verify(sk, zero);
|
||||
}
|
||||
|
||||
/// @brief Weight-1 sketch gated on extractable (2,3) keys.
|
||||
/// @throws std::invalid_argument if any key lacks `dpf::extractable`
|
||||
template <typename K1, typename K2, typename K3, typename RRange>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
bool sketch_verify3(const K1 & k1, const K2 & k2, const K3 & k3,
|
||||
const std::vector<fp61> & s1, const std::vector<fp61> & s2,
|
||||
const std::vector<fp61> & s3, RRange && challenges)
|
||||
{
|
||||
static_assert(K1::is_dpf3 && K2::is_dpf3 && K3::is_dpf3, "dpf3 keys");
|
||||
if (!k1.extractable || !k2.extractable || !k3.extractable)
|
||||
throw std::invalid_argument(
|
||||
"sketch_verify3: keys were not generated with dpf::extractable");
|
||||
return sketch_verify3(s1, s2, s3, std::forward<RRange>(challenges));
|
||||
}
|
||||
|
||||
/// @brief Point proofs plus weight-1 on the opened Shamir full-domain vector.
|
||||
/// @details Completes the paper's three-party statistic after inner verifies.
|
||||
/// Ungated; prefer the key-taking overload for extractable keys.
|
||||
template <typename RRange>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
bool verify_dpf3(const dpf3_proof & p1, const dpf3_proof & p2,
|
||||
const dpf3_proof & p3, const std::vector<fp61> & s1,
|
||||
const std::vector<fp61> & s2, const std::vector<fp61> & s3,
|
||||
RRange && challenges)
|
||||
{
|
||||
if (!verify_dpf3(p1, p2, p3))
|
||||
return false;
|
||||
return sketch_verify3(s1, s2, s3, std::forward<RRange>(challenges));
|
||||
}
|
||||
|
||||
/// @brief Verifiable + extractable check: proofs then gated weight-1 sketch.
|
||||
/// @throws std::invalid_argument if any key lacks `dpf::extractable`
|
||||
template <typename K1, typename K2, typename K3, typename RRange>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
bool verify_dpf3(const K1 & k1, const K2 & k2, const K3 & k3,
|
||||
const dpf3_proof & p1, const dpf3_proof & p2, const dpf3_proof & p3,
|
||||
const std::vector<fp61> & s1, const std::vector<fp61> & s2,
|
||||
const std::vector<fp61> & s3, RRange && challenges)
|
||||
{
|
||||
static_assert(K1::is_dpf3 && K2::is_dpf3 && K3::is_dpf3, "dpf3 keys");
|
||||
if (!k1.extractable || !k2.extractable || !k3.extractable)
|
||||
throw std::invalid_argument(
|
||||
"verify_dpf3: keys were not generated with dpf::extractable");
|
||||
if (!k1.verifiable || !k2.verifiable || !k3.verifiable)
|
||||
throw std::invalid_argument(
|
||||
"verify_dpf3: keys were not generated with dpf::verifiable");
|
||||
return verify_dpf3(p1, p2, p3, s1, s2, s3,
|
||||
std::forward<RRange>(challenges));
|
||||
}
|
||||
|
||||
/// @brief Fresh three-party keys at the same point (new trees — not an update).
|
||||
/// @details Same tags as `make_dpf3`. Use when the key was not generated
|
||||
/// `updatable`, or when the dealer chooses to re-key. Moving `α`
|
||||
/// also requires this path.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
typename ...Tags>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto remake_dpf3(InputT alpha, fp61 beta_new, Tags ...tags)
|
||||
{
|
||||
return make_dpf3<InteriorPRG, ExteriorPRG>(alpha, beta_new, tags...);
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_DPF3_HPP__
|
||||
458
include/dpf/dpf3_cmp.hpp
Normal file
458
include/dpf/dpf3_cmp.hpp
Normal file
|
|
@ -0,0 +1,458 @@
|
|||
/// @file dpf/dpf3_cmp.hpp
|
||||
/// @brief Three-evaluator comparison, blocked comparison, and interval keys.
|
||||
/// @details Each evaluator holds one two-party DCF share of a shared tree
|
||||
/// (party 1 and 3 hold the party-0 half; party 2 holds the party-1
|
||||
/// half). Evaluating a single half yields that party's complementary
|
||||
/// share of the predicate; open with `reconstruct_cmp_halves` on any
|
||||
/// authorized pair (1+2 or 2+3). The threshold is never assembled by
|
||||
/// locally reconstructing a full two-party key. Interval containment
|
||||
/// reuses the three-party comparison; `lo`/`hi` stay public as in F_IC.
|
||||
/// @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_DPF3_CMP_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_DPF3_CMP_HPP__
|
||||
|
||||
#include <cstdint>
|
||||
#include <limits>
|
||||
#include <stdexcept>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/dcf.hpp"
|
||||
#include "dpf/eval_full.hpp"
|
||||
#include "dpf/eval_interval.hpp"
|
||||
#include "dpf/eval_sequence.hpp"
|
||||
#include "dpf/eval_unified.hpp"
|
||||
#include "dpf/fp61.hpp"
|
||||
#include "dpf/incremental.hpp"
|
||||
#include "dpf/interval.hpp"
|
||||
#include "dpf/path_memoizer.hpp"
|
||||
#include "dpf/placement.hpp"
|
||||
#include "dpf/prg_aes.hpp"
|
||||
#include "dpf/shamir3.hpp"
|
||||
#include "dpf/wildcard.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
/// @brief One party's comparison key: a single DCF share plus a Shamir payload tip.
|
||||
/// @tparam Party party index in `{1, 2, 3}`
|
||||
/// @tparam Key two-party DCF key type (party-0 or party-1 half)
|
||||
template <int Party, typename Key>
|
||||
struct dpf3_cmp_key
|
||||
{
|
||||
static_assert(Party >= 1 && Party <= 3, "dpf3_cmp party is 1, 2, or 3");
|
||||
static constexpr int party = Party;
|
||||
static constexpr bool is_dpf3_cmp = true;
|
||||
static constexpr bool is_dpf3 = false;
|
||||
|
||||
using input_type = typename Key::input_type;
|
||||
using key_type = Key;
|
||||
|
||||
/// @brief Inner comparison half. Eval that accepts a `dpf_key` also accepts
|
||||
/// this object and reads `dpf_key`.
|
||||
Key dpf_key{};
|
||||
shamir3::share beta_share{};
|
||||
};
|
||||
|
||||
namespace detail
|
||||
{
|
||||
namespace dpf3_cmp_impl
|
||||
{
|
||||
|
||||
template <typename K0, typename K1>
|
||||
auto assign_share_pair(K0 k0, K1 k1, shamir3::share sh, uint64_t if_false_u)
|
||||
{
|
||||
const uint64_t mask = k0.cmp().mask;
|
||||
const uint64_t di = sh.value.raw() & mask;
|
||||
const uint64_t target =
|
||||
detail::incr::cmp_assign_target(k0.cmp(), di, if_false_u & mask);
|
||||
uint64_t add0 = 0, add1 = 0;
|
||||
const uint64_t r = detail::dcf_impl::sample_addend_blind(mask,
|
||||
[] { return dpf::uniform_sample<typename K0::interior_node>(); });
|
||||
detail::incr::split_cmp_addend(target, mask, r, add0, add1);
|
||||
assign_cmp_local(k0, di, add0);
|
||||
assign_cmp_local(k1, di, add1);
|
||||
return std::make_tuple(std::move(k0), std::move(k1), sh);
|
||||
}
|
||||
|
||||
template <typename K0, typename K1>
|
||||
auto wrap_halves(K0 half0, K1 half1, const std::array<shamir3::share, 3> & shares)
|
||||
{
|
||||
K0 half0_copy = half0; // party 3 holds the same half as party 1
|
||||
dpf3_cmp_key<1, K0> out1{std::move(half0), shares[0]};
|
||||
dpf3_cmp_key<2, K1> out2{std::move(half1), shares[1]};
|
||||
dpf3_cmp_key<3, K0> out3{std::move(half0_copy), shares[2]};
|
||||
return std::make_tuple(std::move(out1), std::move(out2), std::move(out3));
|
||||
}
|
||||
|
||||
/// @brief Build three keys that each hold one DCF half of a shared tree.
|
||||
/// @details One wild comparison tree; δ is planted once into the public value
|
||||
/// words. Absorb is split once. Each evaluator keeps a single half
|
||||
/// (Fig-3 style overlap): party 1 → k0, party 2 → k1, party 3 → k0.
|
||||
/// Eval of one half is that party's complementary share; open with
|
||||
/// `reconstruct_cmp_halves`. A Shamir tip tracks the payload for
|
||||
/// bookkeeping / updates.
|
||||
template <typename InteriorPRG, typename ExteriorPRG, typename InputT,
|
||||
typename Spec>
|
||||
auto make_from_spec(InputT thresh, Spec spec, uint64_t if_true_u,
|
||||
uint64_t if_false_u)
|
||||
{
|
||||
using Concrete = uint64_t;
|
||||
auto wild = dpf::make_dpf<InteriorPRG, ExteriorPRG>(thresh, spec);
|
||||
auto base0 = std::move(wild.first);
|
||||
auto base1 = std::move(wild.second);
|
||||
using K0 = decltype(base0);
|
||||
|
||||
const uint64_t mask = base0.cmp().mask;
|
||||
const uint64_t delta =
|
||||
detail::dcf_impl::beta_delta_u64(
|
||||
detail::dcf_impl::u64_to_beta<Concrete>(if_true_u),
|
||||
detail::dcf_impl::u64_to_beta<Concrete>(if_false_u), mask);
|
||||
const auto shares = shamir3::share_secret(fp61{delta});
|
||||
|
||||
// One shared tree; public δ is the clear payload difference. Absorb is
|
||||
// split once. Each evaluator keeps a single half (Fig-3 style overlap):
|
||||
// party 1 → k0, party 2 → k1, party 3 → k0.
|
||||
uint64_t add0 = 0, add1 = 0;
|
||||
const uint64_t target =
|
||||
detail::incr::cmp_assign_target(base0.cmp(), delta, if_false_u & mask);
|
||||
const uint64_t r = detail::dcf_impl::sample_addend_blind(mask,
|
||||
[] { return dpf::uniform_sample<typename K0::interior_node>(); });
|
||||
detail::incr::split_cmp_addend(target, mask, r, add0, add1);
|
||||
assign_cmp_local(base0, delta, add0);
|
||||
assign_cmp_local(base1, delta, add1);
|
||||
return wrap_halves(std::move(base0), std::move(base1), shares);
|
||||
}
|
||||
|
||||
template <typename KeyT, typename Query, typename PathMemoizer>
|
||||
std::uint64_t eval_one(const KeyT & key, Query && x, PathMemoizer && path)
|
||||
{
|
||||
const auto y = dpf::eval_point(dpf::cmp, key.dpf_key, std::forward<Query>(x),
|
||||
std::forward<PathMemoizer>(path));
|
||||
if constexpr (is_secret_share_v<std::decay_t<decltype(y)>>)
|
||||
return static_cast<std::uint64_t>(y.raw());
|
||||
else
|
||||
return static_cast<std::uint64_t>(y);
|
||||
}
|
||||
|
||||
/// @brief Open complementary DCF halves (additive uint64 shares → fp61).
|
||||
inline fp61 reconstruct_cmp_halves(std::uint64_t k0_share, std::uint64_t k1_share)
|
||||
{
|
||||
return fp61{k0_share + k1_share};
|
||||
}
|
||||
|
||||
} // namespace dpf3_cmp_impl
|
||||
} // namespace detail
|
||||
|
||||
/// @brief Open complementary DCF halves (k0-holder with k1-holder).
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline fp61 reconstruct_cmp_halves(std::uint64_t k0_share, std::uint64_t k1_share)
|
||||
{
|
||||
return detail::dpf3_cmp_impl::reconstruct_cmp_halves(k0_share, k1_share);
|
||||
}
|
||||
|
||||
/// @brief Generate three `lt` comparison keys with a Shamir-shared payload.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto make_dpf3_cmp(InputT thresh, uint64_t if_true, uint64_t if_false = 0)
|
||||
{
|
||||
return detail::dpf3_cmp_impl::make_from_spec<InteriorPRG, ExteriorPRG>(
|
||||
thresh, dpf::lt(dpf::wildcard_value<uint64_t>{}), if_true, if_false);
|
||||
}
|
||||
|
||||
/// @brief Generate three comparison keys from any comparison spec.
|
||||
/// @details Accepts the same packs as two-party `make_dpf`: `lt`/`leq`/`gt`/`geq`,
|
||||
/// `*_at`, `idcf`, `block_width`, and path paints. The shape is kept;
|
||||
/// the payload is installed on a wildcard channel so `update_payload_cmp`
|
||||
/// can rewrite it. A spec that is already wildcard is left unassigned.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
typename Spec,
|
||||
typename = std::enable_if_t<is_cmp_spec_v<std::decay_t<Spec>>>>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto make_dpf3_cmp(InputT thresh, Spec spec)
|
||||
{
|
||||
using S = std::decay_t<Spec>;
|
||||
using Beta = typename S::beta_type;
|
||||
using Concrete = dpf::concrete_type_t<Beta>;
|
||||
if constexpr (is_wildcard_v<Beta>)
|
||||
{
|
||||
auto keys = dpf::make_dpf<InteriorPRG, ExteriorPRG>(
|
||||
thresh, std::move(spec));
|
||||
const auto shares = shamir3::share_secret(fp61{0});
|
||||
return detail::dpf3_cmp_impl::wrap_halves(
|
||||
std::move(keys.first), std::move(keys.second), shares);
|
||||
}
|
||||
else
|
||||
{
|
||||
static_assert(utils::bitlength_of_v<Concrete> <= 64
|
||||
&& !detail::has_from_seed<Concrete>::value,
|
||||
"dpf3 comparison payloads use the uint64 ring");
|
||||
const uint64_t t = detail::dcf_impl::beta_to_u64_simple(spec.if_true,
|
||||
~uint64_t{0});
|
||||
const uint64_t f = detail::dcf_impl::beta_to_u64_simple(spec.if_false,
|
||||
~uint64_t{0});
|
||||
return detail::dpf3_cmp_impl::make_from_spec<InteriorPRG, ExteriorPRG>(
|
||||
thresh, cmp_spec_as_wildcard(std::move(spec)), t, f);
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Blocked comparison with Shamir-shared value words.
|
||||
template <std::size_t B,
|
||||
typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto make_dpf3_cmp_blocked(InputT thresh, uint64_t if_true,
|
||||
uint64_t if_false = 0)
|
||||
{
|
||||
return detail::dpf3_cmp_impl::make_from_spec<InteriorPRG, ExteriorPRG>(
|
||||
thresh, dpf::block_width<B>(dpf::lt(dpf::wildcard_value<uint64_t>{})),
|
||||
if_true, if_false);
|
||||
}
|
||||
|
||||
/// @brief Evaluate a three-party comparison key at `x` (one DCF share).
|
||||
/// @details Returns an unreduced additive `uint64` share. Open a complementary
|
||||
/// pair with `reconstruct_cmp_halves` (sum, then reduce into `fp61`).
|
||||
template <typename KeyT, typename Query,
|
||||
typename PathMemoizer = basic_path_memoizer<typename KeyT::key_type>,
|
||||
std::enable_if_t<std::decay_t<KeyT>::is_dpf3_cmp, int> = 0>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::uint64_t eval_dpf3_cmp(const KeyT & key, Query && x,
|
||||
PathMemoizer && path = PathMemoizer{})
|
||||
{
|
||||
return detail::dpf3_cmp_impl::eval_one(key, std::forward<Query>(x),
|
||||
std::forward<PathMemoizer>(path));
|
||||
}
|
||||
|
||||
/// @brief Full-domain comparison into an additive-share buffer (one DCF expand).
|
||||
/// \complexity Same expansion as `eval_interval` on the whole domain. L = 2^{n - lg(outputs_per_leaf)} leaf nodes, n = `depth`. Time Θ(L) interior traversals. The output buffer stores one slot per domain point (2^n).
|
||||
template <typename KeyT,
|
||||
std::enable_if_t<std::decay_t<KeyT>::is_dpf3_cmp, int> = 0>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::vector<std::uint64_t> eval_full(const KeyT & key)
|
||||
{
|
||||
using in = typename KeyT::input_type;
|
||||
const auto from = std::numeric_limits<in>::min();
|
||||
const auto to = std::numeric_limits<in>::max();
|
||||
auto buf = dpf::make_output_buffer(dpf::cmp, key.dpf_key, from, to);
|
||||
dpf::eval_interval(dpf::cmp, key.dpf_key, from, to, buf);
|
||||
std::vector<std::uint64_t> out(buf.size());
|
||||
for (std::size_t i = 0; i < buf.size(); ++i)
|
||||
{
|
||||
if constexpr (is_secret_share_v<std::decay_t<decltype(buf[i])>>)
|
||||
out[i] = static_cast<std::uint64_t>(buf[i].raw());
|
||||
else
|
||||
out[i] = static_cast<std::uint64_t>(buf[i]);
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief Interval comparison into an additive-share buffer.
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <typename KeyT, typename LaneT,
|
||||
std::enable_if_t<std::decay_t<KeyT>::is_dpf3_cmp, int> = 0>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::vector<std::uint64_t> eval_interval(const KeyT & key, LaneT from, LaneT to)
|
||||
{
|
||||
auto buf = dpf::make_output_buffer(dpf::cmp, key.dpf_key, from, to);
|
||||
dpf::eval_interval(dpf::cmp, key.dpf_key, from, to, buf);
|
||||
std::vector<std::uint64_t> out(buf.size());
|
||||
for (std::size_t i = 0; i < buf.size(); ++i)
|
||||
{
|
||||
if constexpr (is_secret_share_v<std::decay_t<decltype(buf[i])>>)
|
||||
out[i] = static_cast<std::uint64_t>(buf[i].raw());
|
||||
else
|
||||
out[i] = static_cast<std::uint64_t>(buf[i]);
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief In-place comparison payload update via the linear value-CW channel.
|
||||
/// @details Applies `assign_cmp_local` with each party's Shamir share of
|
||||
/// `δ' − δ`, adding that increment into the existing value CWs and
|
||||
/// absorb addends. The threshold / spine stays. `if_false` is fixed
|
||||
/// at 0 (the common `lt(β)` case); changing `if_false` needs a remake.
|
||||
template <typename K1, typename K2, typename K3>
|
||||
void update_payload_cmp(K1 & k1, K2 & k2, K3 & k3, uint64_t if_true_old,
|
||||
uint64_t if_true_new)
|
||||
{
|
||||
static_assert(K1::is_dpf3_cmp && K2::is_dpf3_cmp && K3::is_dpf3_cmp,
|
||||
"dpf3_cmp keys");
|
||||
const uint64_t mask = k1.dpf_key.cmp().mask;
|
||||
// Clear δ on the shared tree; difference is taken in fp61.
|
||||
const uint64_t d =
|
||||
(fp61{if_true_new} - fp61{if_true_old}).raw() & mask;
|
||||
const uint64_t target = detail::incr::cmp_assign_target(k1.dpf_key.cmp(), d, 0);
|
||||
uint64_t add0 = 0, add1 = 0;
|
||||
const uint64_t r = detail::dcf_impl::sample_addend_blind(mask,
|
||||
[] {
|
||||
return dpf::uniform_sample<
|
||||
typename std::decay_t<decltype(k1.dpf_key)>::interior_node>();
|
||||
});
|
||||
detail::incr::split_cmp_addend(target, mask, r, add0, add1);
|
||||
auto bump = [&](auto & key, uint64_t add) {
|
||||
const uint64_t old = [&] {
|
||||
if constexpr (is_party_key_v<std::decay_t<decltype(key.dpf_key)>>)
|
||||
return static_cast<uint64_t>(key.dpf_key.cmp_addend().raw());
|
||||
else if constexpr (std::is_integral_v<
|
||||
std::decay_t<decltype(key.dpf_key.cmp_addend())>>)
|
||||
return static_cast<uint64_t>(key.dpf_key.cmp_addend());
|
||||
else
|
||||
return static_cast<uint64_t>(key.dpf_key.cmp_addend().raw());
|
||||
}();
|
||||
assign_cmp_local(key.dpf_key, d, (old + add) & mask);
|
||||
};
|
||||
bump(k1, add0);
|
||||
bump(k2, add1);
|
||||
bump(k3, add0);
|
||||
const auto shares = shamir3::share_secret(
|
||||
fp61{if_true_new} - fp61{if_true_old});
|
||||
k1.beta_share = shamir3::share{k1.party, k1.beta_share.value + shares[0].value};
|
||||
k2.beta_share = shamir3::share{k2.party, k2.beta_share.value + shares[1].value};
|
||||
k3.beta_share = shamir3::share{k3.party, k3.beta_share.value + shares[2].value};
|
||||
}
|
||||
|
||||
/// @brief Fresh comparison keys for a new payload (new spines — not an update).
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto remake_dpf3_cmp(InputT thresh, uint64_t if_true_new,
|
||||
uint64_t if_false_new = 0)
|
||||
{
|
||||
return make_dpf3_cmp<InteriorPRG, ExteriorPRG>(thresh, if_true_new,
|
||||
if_false_new);
|
||||
}
|
||||
|
||||
/// @brief One party's interval key on a three-party comparison.
|
||||
/// @tparam Party party index in `{1, 2, 3}`
|
||||
/// @tparam CmpKey a `dpf3_cmp_key`
|
||||
template <int Party, typename CmpKey>
|
||||
struct dpf3_ic_key
|
||||
{
|
||||
static constexpr int party = Party;
|
||||
static constexpr bool is_dpf3_ic = true;
|
||||
static constexpr bool is_dpf3 = false;
|
||||
using input_type = typename CmpKey::input_type;
|
||||
using cmp_key_type = CmpKey;
|
||||
|
||||
/// @brief Inner three-party comparison. Its `dpf_key` is the DPF half.
|
||||
CmpKey dpf_key{};
|
||||
uint64_t lo = 0; // public bounds (F_IC)
|
||||
uint64_t hi = 0;
|
||||
uint64_t input_mask = 0;
|
||||
uint64_t group_mask = 0;
|
||||
/// Additive half of δ (k0-holder / k1-holder), opened with `reconstruct_cmp_halves`.
|
||||
std::uint64_t delta_share = 0;
|
||||
std::uint64_t cr_share = 0;
|
||||
};
|
||||
|
||||
/// @brief Generate three interval-containment keys.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto make_dpf3_ic(InputT r, InputT p, InputT q, uint64_t if_true,
|
||||
uint64_t if_false = 0)
|
||||
{
|
||||
auto two = dpf::make_dpf(r, dpf::ic(p, q, if_true, if_false));
|
||||
const uint64_t nmask = two.first.input_mask;
|
||||
const uint64_t gmask = two.first.group_mask;
|
||||
const uint64_t lo = static_cast<uint64_t>(p);
|
||||
const uint64_t hi = static_cast<uint64_t>(q);
|
||||
const InputT gamma = detail::ic_impl::gamma_of(r);
|
||||
|
||||
const uint64_t delta = (if_true - if_false) & gmask;
|
||||
auto cmp_keys = make_dpf3_cmp<InteriorPRG, ExteriorPRG>(gamma, delta, 0);
|
||||
|
||||
uint64_t cr0 = 0, cr1 = 0;
|
||||
if constexpr (is_secret_share_v<decltype(two.first.cr_share)>)
|
||||
{
|
||||
cr0 = static_cast<uint64_t>(two.first.cr_share.raw());
|
||||
cr1 = static_cast<uint64_t>(two.second.cr_share.raw());
|
||||
}
|
||||
else
|
||||
{
|
||||
cr0 = static_cast<uint64_t>(two.first.cr_share);
|
||||
cr1 = static_cast<uint64_t>(two.second.cr_share);
|
||||
}
|
||||
const uint64_t cr_clear = (cr0 + cr1) & gmask;
|
||||
// Additive halves so reconstruct_cmp_halves (sum) recovers the clear words.
|
||||
const uint64_t d0 = dpf::uniform_sample<std::uint64_t>() & gmask;
|
||||
const uint64_t d1 = (delta - d0) & gmask;
|
||||
const uint64_t c0 = dpf::uniform_sample<std::uint64_t>() & gmask;
|
||||
const uint64_t c1 = (cr_clear - c0) & gmask;
|
||||
|
||||
using Cmp1 = std::decay_t<decltype(std::get<0>(cmp_keys))>;
|
||||
using Cmp2 = std::decay_t<decltype(std::get<1>(cmp_keys))>;
|
||||
using Cmp3 = std::decay_t<decltype(std::get<2>(cmp_keys))>;
|
||||
|
||||
dpf3_ic_key<1, Cmp1> k1{std::move(std::get<0>(cmp_keys)), lo, hi, nmask,
|
||||
gmask, d0, c0};
|
||||
dpf3_ic_key<2, Cmp2> k2{std::move(std::get<1>(cmp_keys)), lo, hi, nmask,
|
||||
gmask, d1, c1};
|
||||
dpf3_ic_key<3, Cmp3> k3{std::move(std::get<2>(cmp_keys)), lo, hi, nmask,
|
||||
gmask, d0, c0};
|
||||
return std::make_tuple(std::move(k1), std::move(k2), std::move(k3));
|
||||
}
|
||||
|
||||
/// @brief Evaluate a three-party interval key at `x`.
|
||||
template <typename KeyT, typename Query,
|
||||
std::enable_if_t<std::decay_t<KeyT>::is_dpf3_ic, int> = 0>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::uint64_t eval_dpf3_ic(const KeyT & key, Query && x)
|
||||
{
|
||||
using in_type = typename KeyT::input_type;
|
||||
const uint64_t xu = detail::ic_impl::bits_of(in_type(x));
|
||||
const uint64_t xp = detail::ic_impl::shift_p(xu, key.lo, key.input_mask);
|
||||
const uint64_t xq = detail::ic_impl::shift_q0(xu, key.hi, key.input_mask);
|
||||
basic_path_memoizer<typename KeyT::cmp_key_type::key_type> path;
|
||||
const uint64_t a = eval_dpf3_cmp(key.dpf_key,
|
||||
detail::ic_impl::input_from_bits<in_type>(xp), path);
|
||||
const uint64_t b = eval_dpf3_cmp(key.dpf_key,
|
||||
detail::ic_impl::input_from_bits<in_type>(xq), path);
|
||||
const int cx = detail::ic_impl::public_cx(xu, key.lo, key.hi, key.input_mask);
|
||||
uint64_t scaled = 0;
|
||||
if (cx == 1)
|
||||
scaled = key.delta_share;
|
||||
else if (cx == -1)
|
||||
scaled = static_cast<uint64_t>(0) - key.delta_share;
|
||||
return (static_cast<uint64_t>(0) - a) + b + key.cr_share + scaled;
|
||||
}
|
||||
|
||||
/// @brief `eval_point` overload for three-party comparison keys.
|
||||
/// \complexity O(n) time. n is `depth`. One interior traversal per level from the memoizer resume index through the leaf. Extra space is the path memoizer (O(n) nodes, or one node if it does not memoize). A proof token adds one fold per level walked.
|
||||
template <typename KeyT, typename Query,
|
||||
typename PathMemoizer = basic_path_memoizer<typename KeyT::key_type>,
|
||||
std::enable_if_t<std::decay_t<KeyT>::is_dpf3_cmp, int> = 0>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::uint64_t eval_point(const KeyT & key, Query && x,
|
||||
PathMemoizer && path = PathMemoizer{})
|
||||
{
|
||||
return eval_dpf3_cmp(key, std::forward<Query>(x),
|
||||
std::forward<PathMemoizer>(path));
|
||||
}
|
||||
|
||||
/// @brief `eval_point` overload for three-party interval keys.
|
||||
/// \complexity O(n) time. n is `depth`. One interior traversal per level from the memoizer resume index through the leaf. Extra space is the path memoizer (O(n) nodes, or one node if it does not memoize). A proof token adds one fold per level walked.
|
||||
template <typename KeyT, typename Query,
|
||||
std::enable_if_t<std::decay_t<KeyT>::is_dpf3_ic, int> = 0>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::uint64_t eval_point(const KeyT & key, Query && x)
|
||||
{
|
||||
return eval_dpf3_ic(key, std::forward<Query>(x));
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_DPF3_CMP_HPP__
|
||||
125
include/dpf/dpf3_ds.hpp
Normal file
125
include/dpf/dpf3_ds.hpp
Normal file
|
|
@ -0,0 +1,125 @@
|
|||
/// @file dpf/dpf3_ds.hpp
|
||||
/// @brief Dual-spine Doerner–Shelat generation of (2,3) point keys.
|
||||
/// @note Two independent openings, one per spine of Zyskind, Yanai, and Pentland (ePrint 2024/1658, Figure 3). Each opening follows Doerner and shelat, CCS 2017 (ePrint 2017/827).
|
||||
/// @details Runs two independent two-party DS walks (spines A and B) with
|
||||
/// Fig-3 `τ` payloads, then assembles the three evaluator keys via
|
||||
/// `assemble_from_spines`. Matches honest-dealer `make_dpf3` on the
|
||||
/// opened point `x0 ⊕ x1` (after the signed-MSB flip on share 0).
|
||||
/// @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_DPF3_DS_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_DPF3_DS_HPP__
|
||||
|
||||
#include <tuple>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/dpf3.hpp"
|
||||
#include "dpf/fp61.hpp"
|
||||
#include "dpf/incremental.hpp"
|
||||
#include "dpf/placement.hpp"
|
||||
#include "dpf/prg_aes.hpp"
|
||||
#include "dpf/verifiable.hpp"
|
||||
#include "dpf/wildcard.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace detail
|
||||
{
|
||||
namespace dpf3_impl
|
||||
{
|
||||
|
||||
template <typename InteriorPRG, typename ExteriorPRG, typename InputT,
|
||||
bool Verifiable, bool Updatable, bool Extractable>
|
||||
auto make_point3_ds(InputT x0, InputT x1, fp61 beta)
|
||||
{
|
||||
using X = xor61;
|
||||
const InputT alpha = open_xor_point(x0, x1);
|
||||
const tau_quad t = sample_taus(beta);
|
||||
const X payload_a = t.t0 + t.t1;
|
||||
const X payload_b = t.t2 + t.t3;
|
||||
|
||||
if constexpr (Updatable)
|
||||
{
|
||||
// Classic DS rejects wildcards; the incremental path plants beaver
|
||||
// leaves. Inner verifiable tags match the outer `Verifiable` flag.
|
||||
auto A = [&] {
|
||||
if constexpr (Verifiable)
|
||||
return dpf::make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
|
||||
x0, x1, dpf::wildcard_value<X>{}, dpf::verifiable{});
|
||||
else
|
||||
return dpf::make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
|
||||
x0, x1, dpf::wildcard_value<X>{});
|
||||
}();
|
||||
auto B = [&] {
|
||||
if constexpr (Verifiable)
|
||||
return dpf::make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
|
||||
x0, x1, dpf::wildcard_value<X>{}, dpf::verifiable{});
|
||||
else
|
||||
return dpf::make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
|
||||
x0, x1, dpf::wildcard_value<X>{});
|
||||
}();
|
||||
assign_xor_payload(A.first, A.second, payload_a);
|
||||
assign_xor_payload(B.first, B.second, payload_b);
|
||||
return assemble_from_spines<Verifiable, Extractable, true>(
|
||||
std::move(A.first), std::move(A.second), std::move(B.first),
|
||||
std::move(B.second), t, alpha);
|
||||
}
|
||||
else
|
||||
{
|
||||
auto A = [&] {
|
||||
if constexpr (Verifiable)
|
||||
return dpf::make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
|
||||
x0, x1, payload_a, dpf::verifiable{});
|
||||
else
|
||||
return dpf::make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
|
||||
x0, x1, payload_a);
|
||||
}();
|
||||
auto B = [&] {
|
||||
if constexpr (Verifiable)
|
||||
return dpf::make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
|
||||
x0, x1, payload_b, dpf::verifiable{});
|
||||
else
|
||||
return dpf::make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
|
||||
x0, x1, payload_b);
|
||||
}();
|
||||
return assemble_from_spines<Verifiable, Extractable, false>(
|
||||
std::move(A.first), std::move(A.second), std::move(B.first),
|
||||
std::move(B.second), t, alpha);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace dpf3_impl
|
||||
} // namespace detail
|
||||
|
||||
/// @brief Dual-spine Doerner–Shelat (2,3) keys for XOR shares of `α`.
|
||||
/// @details `α = x0 ⊕ x1` after the signed-MSB flip on `x0`. Tags match
|
||||
/// `make_dpf3`: `verifiable`, `extractable`, and `updatable`, any order.
|
||||
/// \complexity O(n) time. One `ds_advance_level` per level: two PRG expansions and one `prepare_level`. n is `depth`.
|
||||
/// \rounds No sockets. This is the in-process transcript. A networked walk is `dpf::party::dist::point_party`.
|
||||
/// \communication none here. `local_cw_protocol` opens the correction word locally.
|
||||
/// \preprocessing Per level, `prepare_level` draws one `ds_cw_pads` (two parties × a 128-bit rand, a 128-bit gamma, and a bit) and two `ds_and_pads`. Arithmetic inputs also run a carry chain of n-1 bit-AND triples in `encode_walk_shares`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
typename ...Tags>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto make_dpf3_doerner_shelat(InputT x0, InputT x1, fp61 beta, Tags ...tags)
|
||||
{
|
||||
using flags = detail::dpf3_impl::tag_flags<Tags...>;
|
||||
static_assert(flags::known,
|
||||
"make_dpf3_doerner_shelat tags are verifiable, extractable, updatable");
|
||||
static_assert(sizeof...(Tags) == flags::counted,
|
||||
"make_dpf3_doerner_shelat: repeated tag");
|
||||
(void)std::initializer_list<int>{((void)tags, 0)...};
|
||||
return detail::dpf3_impl::make_point3_ds<InteriorPRG, ExteriorPRG, InputT,
|
||||
flags::verifiable, flags::updatable, flags::extractable>(x0, x1, beta);
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_DPF3_DS_HPP__
|
||||
336
include/dpf/dpf3_multipoint.hpp
Normal file
336
include/dpf/dpf3_multipoint.hpp
Normal file
|
|
@ -0,0 +1,336 @@
|
|||
/// @file dpf/dpf3_multipoint.hpp
|
||||
/// @brief Three-evaluator multipoint DPF: Figure 3 of Zyskind, Yanai, and Pentland (ePrint 2024/1658) per bucket.
|
||||
/// @details Cuckoo packing matches `make_multipoint`. Each bucket is a (2,3)
|
||||
/// point key. Evaluation sums Shamir shares across the three probes.
|
||||
/// Updatable multipoint keys keep beaver leaves on every bucket;
|
||||
/// `update_payload` replays `insert_cuckoo` with the existing `σ`
|
||||
/// (deterministic) and runs Fig-10 per occupied bucket — no re-cuckoo.
|
||||
/// @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_DPF3_MULTIPOINT_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_DPF3_MULTIPOINT_HPP__
|
||||
|
||||
#include <algorithm>
|
||||
#include <cstdint>
|
||||
#include <stdexcept>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
#include "simde/simde/x86/avx2.h"
|
||||
|
||||
#include "dpf/dpf3.hpp"
|
||||
#include "dpf/fp61.hpp"
|
||||
#include "dpf/multipoint.hpp"
|
||||
#include "dpf/path_memoizer.hpp"
|
||||
#include "dpf/prg_aes.hpp"
|
||||
#include "dpf/random.hpp"
|
||||
#include "dpf/shamir3.hpp"
|
||||
#include "dpf/verifiable.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
/// @brief One party's multipoint (2,3) key.
|
||||
/// @tparam Party party index in `{1, 2, 3}`
|
||||
/// @tparam InputT input domain type
|
||||
/// @tparam BucketKey a `dpf3_key` for that party
|
||||
template <int Party, typename InputT, typename BucketKey>
|
||||
struct multipoint3_key
|
||||
{
|
||||
static_assert(Party >= 1 && Party <= 3, "multipoint3 party is 1, 2, or 3");
|
||||
static constexpr int party = Party;
|
||||
static constexpr bool is_multipoint3 = true;
|
||||
static constexpr bool is_dpf3 = false;
|
||||
static constexpr std::size_t kappa = 3;
|
||||
|
||||
using input_type = InputT;
|
||||
using bucket_key = BucketKey;
|
||||
using bucket_input = typename BucketKey::input_type;
|
||||
|
||||
simde__m128i sigma{};
|
||||
std::uint64_t bucket_count = 0;
|
||||
mpf_word bucket_domain{};
|
||||
std::vector<BucketKey> buckets{};
|
||||
bool verifiable = false;
|
||||
bool extractable = false;
|
||||
bool updatable = false;
|
||||
};
|
||||
|
||||
namespace detail
|
||||
{
|
||||
namespace mpf3
|
||||
{
|
||||
|
||||
template <bool Verifiable, bool Extractable, bool Updatable,
|
||||
typename InteriorPRG, typename ExteriorPRG, typename InputT>
|
||||
auto make_impl(std::vector<InputT> alphas, std::vector<fp61> betas,
|
||||
multipoint_params params)
|
||||
{
|
||||
namespace mpf = detail::mpf;
|
||||
if (alphas.size() != betas.size())
|
||||
throw std::invalid_argument("make_multipoint3: point/payload count");
|
||||
if (alphas.empty())
|
||||
throw std::invalid_argument("make_multipoint3: no points");
|
||||
{
|
||||
auto sorted = alphas;
|
||||
std::sort(sorted.begin(), sorted.end());
|
||||
if (std::adjacent_find(sorted.begin(), sorted.end()) != sorted.end())
|
||||
throw std::invalid_argument("make_multipoint3: duplicate points");
|
||||
}
|
||||
|
||||
const auto t = static_cast<std::uint64_t>(alphas.size());
|
||||
const auto m = mpf::bucket_count_for(t, params.lambda);
|
||||
constexpr std::size_t input_bits = utils::bitlength_of_v<InputT>;
|
||||
const mpf_word n = mpf::domain_size(input_bits);
|
||||
constexpr int kappa = 3;
|
||||
const mpf_word span = mpf::word_mul_small(n, kappa);
|
||||
const mpf_word den{uint256_t{m}, uint256_t{0}};
|
||||
const mpf_word numer = mpf::word_add(span,
|
||||
mpf::word_sub(den, mpf_word{uint256_t{1}, uint256_t{0}}));
|
||||
const mpf_word b = mpf::word_divmod(numer, den).first;
|
||||
|
||||
const int attempts = params.retries < 1 ? 1 : params.retries;
|
||||
for (int attempt = 0; attempt < attempts; ++attempt)
|
||||
{
|
||||
try
|
||||
{
|
||||
const simde__m128i sigma = dpf::uniform_sample<simde__m128i>();
|
||||
std::vector<mpf::slot> table;
|
||||
if (!mpf::insert_cuckoo(sigma, alphas, m, n, b, params.max_evictions,
|
||||
table))
|
||||
continue;
|
||||
|
||||
// Probe the first bucket type to name the three party key types.
|
||||
InputT gamma0{};
|
||||
fp61 beta0{};
|
||||
if (table[0].item >= 0)
|
||||
{
|
||||
const auto & alpha = alphas[static_cast<std::size_t>(table[0].item)];
|
||||
const auto loc = mpf::locate(sigma, mpf::to_word(alpha),
|
||||
table[0].hash, n, b);
|
||||
gamma0 = mpf::from_word<InputT>(loc.index);
|
||||
beta0 = betas[static_cast<std::size_t>(table[0].item)];
|
||||
}
|
||||
auto make_bucket3 = [&](InputT gamma, fp61 beta) {
|
||||
return detail::dpf3_impl::make_tagged<Verifiable, Extractable,
|
||||
Updatable, InteriorPRG, ExteriorPRG>(gamma, beta);
|
||||
};
|
||||
auto first = make_bucket3(gamma0, beta0);
|
||||
using K1 = std::decay_t<decltype(std::get<0>(first))>;
|
||||
using K2 = std::decay_t<decltype(std::get<1>(first))>;
|
||||
using K3 = std::decay_t<decltype(std::get<2>(first))>;
|
||||
|
||||
multipoint3_key<1, InputT, K1> left;
|
||||
multipoint3_key<2, InputT, K2> mid;
|
||||
multipoint3_key<3, InputT, K3> right;
|
||||
left.sigma = mid.sigma = right.sigma = sigma;
|
||||
left.bucket_count = mid.bucket_count = right.bucket_count = m;
|
||||
left.bucket_domain = mid.bucket_domain = right.bucket_domain = b;
|
||||
left.verifiable = mid.verifiable = right.verifiable = Verifiable;
|
||||
left.extractable = mid.extractable = right.extractable = Extractable;
|
||||
left.updatable = mid.updatable = right.updatable = Updatable;
|
||||
left.buckets.reserve(static_cast<std::size_t>(m));
|
||||
mid.buckets.reserve(static_cast<std::size_t>(m));
|
||||
right.buckets.reserve(static_cast<std::size_t>(m));
|
||||
|
||||
left.buckets.push_back(std::move(std::get<0>(first)));
|
||||
mid.buckets.push_back(std::move(std::get<1>(first)));
|
||||
right.buckets.push_back(std::move(std::get<2>(first)));
|
||||
|
||||
for (std::uint64_t i = 1; i < m; ++i)
|
||||
{
|
||||
InputT gamma{};
|
||||
fp61 beta{};
|
||||
if (table[static_cast<std::size_t>(i)].item >= 0)
|
||||
{
|
||||
const auto & alpha = alphas[static_cast<std::size_t>(
|
||||
table[static_cast<std::size_t>(i)].item)];
|
||||
const auto loc = mpf::locate(sigma, mpf::to_word(alpha),
|
||||
table[static_cast<std::size_t>(i)].hash, n, b);
|
||||
if (loc.bucket != i)
|
||||
throw mpf::prp_walk_error{};
|
||||
gamma = mpf::from_word<InputT>(loc.index);
|
||||
beta = betas[static_cast<std::size_t>(
|
||||
table[static_cast<std::size_t>(i)].item)];
|
||||
}
|
||||
auto made = make_bucket3(gamma, beta);
|
||||
left.buckets.push_back(std::move(std::get<0>(made)));
|
||||
mid.buckets.push_back(std::move(std::get<1>(made)));
|
||||
right.buckets.push_back(std::move(std::get<2>(made)));
|
||||
}
|
||||
return std::make_tuple(std::move(left), std::move(mid),
|
||||
std::move(right));
|
||||
}
|
||||
catch (const mpf::prp_walk_error &)
|
||||
{
|
||||
continue;
|
||||
}
|
||||
}
|
||||
throw std::runtime_error("make_multipoint3: cuckoo hashing failed");
|
||||
}
|
||||
|
||||
template <typename Key>
|
||||
fp61 eval_at(const Key & key, typename Key::input_type x)
|
||||
{
|
||||
namespace mpf = detail::mpf;
|
||||
constexpr std::size_t input_bits = utils::bitlength_of_v<typename Key::input_type>;
|
||||
const mpf_word n = mpf::domain_size(input_bits);
|
||||
const mpf_word b = key.bucket_domain;
|
||||
fp61 sum{};
|
||||
using InnerA = typename Key::bucket_key::plus_a_type::inner_type;
|
||||
using InnerB = typename Key::bucket_key::plus_b_type::inner_type;
|
||||
basic_path_memoizer<InnerA> path_a{};
|
||||
basic_path_memoizer<InnerB> path_b{};
|
||||
for (int hash = 0; hash < static_cast<int>(Key::kappa); ++hash)
|
||||
{
|
||||
const auto loc = mpf::locate(key.sigma, mpf::to_word(x), hash, n, b);
|
||||
if (loc.bucket >= key.bucket_count)
|
||||
throw std::runtime_error("multipoint3 eval: bucket out of range");
|
||||
const auto gamma = mpf::from_word<typename Key::bucket_input>(loc.index);
|
||||
sum = sum + dpf::eval_point(key.buckets[loc.bucket], gamma, path_a,
|
||||
path_b);
|
||||
}
|
||||
return sum;
|
||||
}
|
||||
|
||||
} // namespace mpf3
|
||||
} // namespace detail
|
||||
|
||||
namespace detail
|
||||
{
|
||||
namespace mpf3
|
||||
{
|
||||
|
||||
template <typename T>
|
||||
struct is_multipoint_params : std::is_same<std::decay_t<T>, multipoint_params>
|
||||
{ };
|
||||
|
||||
template <typename ...Ts>
|
||||
using last_type_t = std::tuple_element_t<sizeof...(Ts) == 0 ? 0 : sizeof...(Ts) - 1,
|
||||
std::tuple<std::decay_t<Ts>..., multipoint_params>>;
|
||||
|
||||
template <typename Run, typename Tuple, std::size_t ...I>
|
||||
auto run_without_last(Run && run, Tuple && pack, std::index_sequence<I...>)
|
||||
{
|
||||
return std::forward<Run>(run)(std::get<sizeof...(I)>(pack),
|
||||
std::get<I>(std::forward<Tuple>(pack))...);
|
||||
}
|
||||
|
||||
} // namespace mpf3
|
||||
} // namespace detail
|
||||
|
||||
/// @brief Cuckoo-pack points into (2,3) point-key buckets.
|
||||
/// @details Tags match `make_dpf3` (`verifiable`, `extractable`, `updatable`),
|
||||
/// in any order. An optional `multipoint_params` may follow the tags.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename AlphaRange,
|
||||
typename BetaRange,
|
||||
typename ...Tail>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto make_multipoint3(const AlphaRange & alphas, const BetaRange & betas,
|
||||
Tail && ...tail)
|
||||
{
|
||||
using input_type = std::decay_t<decltype(*std::begin(alphas))>;
|
||||
auto run = [&](multipoint_params params, auto && ...tags) {
|
||||
using flags = detail::dpf3_impl::tag_flags<std::decay_t<decltype(tags)>...>;
|
||||
static_assert(flags::known,
|
||||
"make_multipoint3 tags are verifiable, extractable, updatable");
|
||||
static_assert(sizeof...(tags) == flags::counted,
|
||||
"make_multipoint3: repeated tag");
|
||||
return detail::mpf3::make_impl<flags::verifiable, flags::extractable,
|
||||
flags::updatable, InteriorPRG, ExteriorPRG>(
|
||||
std::vector<input_type>(std::begin(alphas), std::end(alphas)),
|
||||
std::vector<fp61>(std::begin(betas), std::end(betas)), params);
|
||||
};
|
||||
if constexpr (sizeof...(Tail) > 0
|
||||
&& detail::mpf3::is_multipoint_params<
|
||||
detail::mpf3::last_type_t<Tail...>>::value)
|
||||
{
|
||||
return detail::mpf3::run_without_last(run,
|
||||
std::forward_as_tuple(std::forward<Tail>(tail)...),
|
||||
std::make_index_sequence<sizeof...(Tail) - 1>{});
|
||||
}
|
||||
else
|
||||
return run(multipoint_params{}, std::forward<Tail>(tail)...);
|
||||
}
|
||||
|
||||
/// @brief Sum the three bucket Shamir shares at `x`.
|
||||
/// \complexity Three `eval_point` calls (`Key::kappa` is 3), each O(n_b) where n_b is the bucket key depth.
|
||||
template <typename Key,
|
||||
std::enable_if_t<Key::is_multipoint3, int> = 0>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
fp61 eval_multipoint(const Key & key, typename Key::input_type x)
|
||||
{
|
||||
return detail::mpf3::eval_at(key, x);
|
||||
}
|
||||
|
||||
/// @brief In-place multipoint payload update; keeps `σ` and cuckoo placement.
|
||||
/// @details Replays `insert_cuckoo` with the existing `σ` (RNG is seeded from
|
||||
/// `σ`, so placement is deterministic) to rediscover each point's
|
||||
/// bucket, then runs Fig-10 `update_payload` on that bucket. Empty
|
||||
/// buckets stay a share of 0. Requires `updatable` generation.
|
||||
/// Point set and order must match generation.
|
||||
template <typename M1, typename M2, typename M3, typename AlphaRange,
|
||||
typename BetaRange>
|
||||
void update_payload(M1 & k1, M2 & k2, M3 & k3, const AlphaRange & alphas,
|
||||
const BetaRange & betas_new, multipoint_params params = {})
|
||||
{
|
||||
static_assert(M1::is_multipoint3 && M2::is_multipoint3
|
||||
&& M3::is_multipoint3,
|
||||
"multipoint3 keys");
|
||||
if (!k1.updatable || !k2.updatable || !k3.updatable)
|
||||
throw std::invalid_argument(
|
||||
"update_payload: multipoint3 keys were not generated updatable");
|
||||
namespace mpf = detail::mpf;
|
||||
using InputT = typename M1::input_type;
|
||||
const auto alpha_v = std::vector<InputT>(std::begin(alphas),
|
||||
std::end(alphas));
|
||||
const auto beta_v = std::vector<fp61>(std::begin(betas_new),
|
||||
std::end(betas_new));
|
||||
if (alpha_v.size() != beta_v.size())
|
||||
throw std::invalid_argument("update_payload: point/payload count");
|
||||
constexpr std::size_t input_bits = utils::bitlength_of_v<InputT>;
|
||||
const mpf_word n = mpf::domain_size(input_bits);
|
||||
const mpf_word b = k1.bucket_domain;
|
||||
std::vector<mpf::slot> table;
|
||||
if (!mpf::insert_cuckoo(k1.sigma, alpha_v, k1.bucket_count, n, b,
|
||||
params.max_evictions, table))
|
||||
throw std::runtime_error(
|
||||
"update_payload: cuckoo replay failed (σ/points mismatch?)");
|
||||
for (std::uint64_t bi = 0; bi < k1.bucket_count; ++bi)
|
||||
{
|
||||
if (table[static_cast<std::size_t>(bi)].item < 0)
|
||||
continue;
|
||||
const auto item = static_cast<std::size_t>(
|
||||
table[static_cast<std::size_t>(bi)].item);
|
||||
const auto loc = mpf::locate(k1.sigma, mpf::to_word(alpha_v[item]),
|
||||
table[static_cast<std::size_t>(bi)].hash, n, b);
|
||||
if (loc.bucket != bi)
|
||||
throw std::runtime_error("update_payload: PRP walk mismatch");
|
||||
const auto gamma =
|
||||
mpf::from_word<typename M1::bucket_input>(loc.index);
|
||||
dpf::update_payload(k1.buckets[bi], k2.buckets[bi], k3.buckets[bi],
|
||||
gamma, beta_v[item]);
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Fresh multipoint3 keys (new cuckoo table — not an in-place update).
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename AlphaRange,
|
||||
typename BetaRange>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto remake_multipoint3(const AlphaRange & alphas, const BetaRange & betas_new,
|
||||
multipoint_params params = {})
|
||||
{
|
||||
return make_multipoint3<InteriorPRG, ExteriorPRG>(alphas, betas_new, params);
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_DPF3_MULTIPOINT_HPP__
|
||||
|
|
@ -17,6 +17,7 @@
|
|||
#include <bitset>
|
||||
#include <atomic>
|
||||
|
||||
#include "dpf/experiment_note.hpp"
|
||||
#include "dpf/prg_aes.hpp"
|
||||
#include "dpf/tree_traits.hpp"
|
||||
#include "dpf/wildcard.hpp"
|
||||
|
|
@ -35,6 +36,12 @@ namespace dpf
|
|||
#ifdef LIBDPF_HAS_ASIO
|
||||
namespace asio
|
||||
{
|
||||
/// @brief Exchange a wildcard output share and install the opened leaf.
|
||||
/// @see dpf::leaf_wrapper
|
||||
/// \complexity O(leaf bytes) for the Beaver leaf arithmetic, plus the transfers.
|
||||
/// \rounds 2. Write/read the blinded output share, then write/read the leaf share (`async_assign_wildcard_output`).
|
||||
/// \communication `sizeof(output_type)` plus `sizeof(leaf_type)` each way.
|
||||
/// \preprocessing The leaf Beaver triple (`vector_blind`, `output_blind`, `blinded_vector`) was stored at keygen.
|
||||
template <std::size_t I,
|
||||
typename PeerT,
|
||||
typename DpfKey,
|
||||
|
|
@ -251,14 +258,14 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
const leaf_tuple & leaves,
|
||||
const beaver_tuple & beavers,
|
||||
input_type offset_share)
|
||||
: root_{root},
|
||||
: leaf_nodes(get_wrappers(leaves, beavers)),
|
||||
offset_x{offset_share},
|
||||
root_{root},
|
||||
correction_words_{correction_words},
|
||||
correction_advice_{correction_advice},
|
||||
mutable_wildcard_mask_{dpf::utils::make_bitset(dpf::is_wildcard_v<OutputT>,
|
||||
dpf::is_wildcard_v<OutputTs>...)},
|
||||
leaf_nodes(get_wrappers(leaves, beavers)),
|
||||
common_part_hash_{utils::get_common_part_hash(correction_words_, correction_advice_, leaf_nodes, wildcard_mask)},
|
||||
offset_x{offset_share}
|
||||
common_part_hash_{utils::get_common_part_hash(correction_words_, correction_advice_, leaf_nodes, wildcard_mask)}
|
||||
{ }
|
||||
classic_dpf_key_impl(const classic_dpf_key_impl &) = default;
|
||||
classic_dpf_key_impl(classic_dpf_key_impl &&) = default;
|
||||
|
|
@ -410,6 +417,29 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
return traverse_exterior<I>(node, std::get<I>(leaf_nodes).get());
|
||||
}
|
||||
|
||||
/// @brief Eight one-block leaves. One `eval_x8` instead of eight `eval` calls.
|
||||
template <std::size_t I = 0, typename Out>
|
||||
void traverse_exterior_x8(const interior_node * HEDLEY_RESTRICT nodes,
|
||||
Out * HEDLEY_RESTRICT out) const noexcept
|
||||
{
|
||||
using output_type = std::tuple_element_t<I, concrete_outputs_tuple>;
|
||||
using block = typename exterior_prg::block_type;
|
||||
alignas(64) block seeds[8];
|
||||
alignas(64) block masks[8];
|
||||
for (std::size_t t = 0; t < 8; ++t)
|
||||
seeds[t] = utils::to_exterior_node<block>(unset_lo_2bits(nodes[t]));
|
||||
constexpr auto pos = dpf::block_offset_of_leaf_v<I, block,
|
||||
concrete_outputs_tuple>;
|
||||
exterior_prg::eval_x8(seeds, masks, static_cast<psnip_uint32_t>(pos));
|
||||
const auto & cw = std::get<I>(leaf_nodes).get();
|
||||
for (std::size_t t = 0; t < 8; ++t)
|
||||
{
|
||||
encode_curve_leaf_mask<concrete_type_t<output_type>>(masks[t]);
|
||||
out[t] = dpf::subtract_leaf<output_type>(
|
||||
dpf::get_if_lo_bit(cw, nodes[t]), masks[t]);
|
||||
}
|
||||
}
|
||||
|
||||
leaf_wrapper_tuple leaf_nodes;
|
||||
offset_type offset_x;
|
||||
static constexpr std::array<bool, sizeof...(OutputTs)+1> wildcard_mask{dpf::is_wildcard_v<OutputT>,
|
||||
|
|
@ -670,6 +700,35 @@ struct cmp_storage
|
|||
}
|
||||
}
|
||||
|
||||
/// @brief `mix(base, coeff)` at every value CW, then install `addend`.
|
||||
/// @tparam Mix word rewriter
|
||||
/// @param mix combines one stored base word with its integer coefficient word
|
||||
/// @param addend this party's share of the constant payload
|
||||
template <typename Mix>
|
||||
void assign_payload(Mix mix, value_cw_word addend)
|
||||
{
|
||||
static_assert(Wild,
|
||||
"assign_cmp on a key whose comparison payload is not a wildcard");
|
||||
if constexpr (Wild)
|
||||
{
|
||||
for (std::size_t i = 0; i < Depth; ++i)
|
||||
value_cw_[i] = mix(value_cw_[i], wild_.value_cw_coeff[i]);
|
||||
if constexpr (Blocked)
|
||||
{
|
||||
for (std::size_t i = 0; i < TailLen; ++i)
|
||||
tail_[i] = mix(tail_[i], wild_.tail_coeff[i]);
|
||||
}
|
||||
cw_last_ = mix(cw_last_, wild_.cw_last_coeff);
|
||||
if constexpr (Idcf)
|
||||
{
|
||||
for (std::size_t i = 0; i < prefix_cw_len; ++i)
|
||||
prefix_cw_[i] = mix(prefix_cw_[i], wild_.prefix_cw_coeff[i]);
|
||||
}
|
||||
cmp_addend_ = addend;
|
||||
wild_.assigned = true;
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Per-level δ coefficients for a wildcard comparison. Empty when the
|
||||
/// payload is concrete.
|
||||
/// @return the coefficient table
|
||||
|
|
@ -767,6 +826,7 @@ struct incr_key_base
|
|||
using interior_node = typename InteriorPRG::block_type;
|
||||
using exterior_node = typename ExteriorPRG::block_type;
|
||||
using input_type = dpf::concrete_type_t<InputT>;
|
||||
using raw_input_type = InputT;
|
||||
using placed_tuple = PlacedTuple;
|
||||
using node_type = exterior_node;
|
||||
static constexpr std::size_t cmp_depth = CmpDepth;
|
||||
|
|
@ -797,10 +857,15 @@ struct incr_key_base
|
|||
static constexpr std::size_t cmp_tail =
|
||||
(CmpBlock == 0 || cmp_q == 0) ? 0 : (std::size_t{1} << cmp_q);
|
||||
/// @brief Narrowest unsigned word that holds `cmp_out_bits` bits (1 byte for a
|
||||
/// bit / ≤8-bit payload, 2 for ≤16, 4 for ≤32, 8 for ≤64). Value CWs and
|
||||
/// the addend share are stored in this word.
|
||||
using value_cw_word = utils::integral_type_from_bitlength_t<
|
||||
(CmpOutBits == 0 ? std::size_t{1} : CmpOutBits)>;
|
||||
/// bit / ≤8-bit payload, 2 for ≤16, 4 for ≤32, 8 for ≤64). Payloads wider
|
||||
/// than 256 bits are stored as their raw bytes. Value CWs and the addend
|
||||
/// share use this word.
|
||||
static constexpr std::size_t cmp_word_bits =
|
||||
CmpOutBits == 0 ? std::size_t{1} : CmpOutBits;
|
||||
using value_cw_word = std::conditional_t<
|
||||
(cmp_word_bits <= 256),
|
||||
utils::integral_type_from_bitlength_t<cmp_word_bits>,
|
||||
std::array<std::uint8_t, (cmp_word_bits + 7) / 8>>;
|
||||
|
||||
static constexpr std::size_t num_outputs = std::tuple_size_v<PlacedTuple>;
|
||||
static constexpr std::size_t input_bits = utils::bitlength_of_v<input_type>;
|
||||
|
|
@ -835,12 +900,15 @@ struct incr_key_base
|
|||
/// @brief Multi-level / comparison keys route through the slot-aware eval path.
|
||||
/// Classic-shaped packs (every slot at full input width, no cmp) keep the
|
||||
/// classic `eval_*` fast path even when verifiable/extractable phantoms are
|
||||
/// present.
|
||||
/// present. `eq` / `eq_at` with a public `if_false` addend also take the
|
||||
/// slot-aware path so party 0 can absorb that addend.
|
||||
static constexpr bool is_multilevel = [] {
|
||||
if constexpr (CmpDepth > 0)
|
||||
return true;
|
||||
if constexpr (num_outputs == 0)
|
||||
return true;
|
||||
else if constexpr (detail::incr::any_public_addend_v<PlacedTuple>)
|
||||
return true;
|
||||
else
|
||||
{
|
||||
for (std::size_t i = 0; i < num_outputs; ++i)
|
||||
|
|
@ -867,6 +935,20 @@ struct incr_key_base
|
|||
template <std::size_t I>
|
||||
using concrete_output_type = concrete_type_t<output_type_t<I>>;
|
||||
|
||||
private:
|
||||
template <std::size_t... Is>
|
||||
static auto outputs_tuple_type(std::index_sequence<Is...>)
|
||||
-> std::tuple<output_type_t<Is>...>;
|
||||
template <std::size_t... Is>
|
||||
static auto concrete_outputs_tuple_type(std::index_sequence<Is...>)
|
||||
-> std::tuple<concrete_output_type<Is>...>;
|
||||
|
||||
public:
|
||||
using outputs_tuple = decltype(outputs_tuple_type(
|
||||
std::make_index_sequence<num_outputs>{}));
|
||||
using concrete_outputs_tuple = decltype(concrete_outputs_tuple_type(
|
||||
std::make_index_sequence<num_outputs>{}));
|
||||
|
||||
template <std::size_t I>
|
||||
static constexpr std::size_t lg_outputs_per_leaf_of =
|
||||
(num_outputs > 0) ? meta[I].lg_opl : 0;
|
||||
|
|
@ -899,6 +981,23 @@ struct incr_key_base
|
|||
static constexpr auto wildcard_mask =
|
||||
wildcard_mask_tuple(std::make_index_sequence<num_outputs>{});
|
||||
|
||||
template <std::size_t... Is>
|
||||
static constexpr std::array<bool, num_outputs>
|
||||
wildcard_mask_array(std::index_sequence<Is...>)
|
||||
{
|
||||
return {{std::get<Is>(wildcard_mask)...}};
|
||||
}
|
||||
static constexpr auto wildcard_bits =
|
||||
wildcard_mask_array(std::make_index_sequence<num_outputs>{});
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_PURE
|
||||
constexpr bool is_wildcard(std::size_t i) const noexcept
|
||||
{
|
||||
return i < num_outputs && wildcard_bits[i];
|
||||
}
|
||||
|
||||
static constexpr std::size_t deepest_prefix = [] {
|
||||
if constexpr (num_outputs == 0)
|
||||
return CmpDepth;
|
||||
|
|
@ -938,6 +1037,14 @@ struct incr_key_base
|
|||
using addend_tuple = decltype(addend_tuple_t(
|
||||
std::make_index_sequence<num_outputs>{}));
|
||||
|
||||
static value_cw_word cw_word_from_u64(std::uint64_t v)
|
||||
{
|
||||
if constexpr (std::is_integral_v<value_cw_word>)
|
||||
return static_cast<value_cw_word>(v);
|
||||
else
|
||||
return value_cw_word{};
|
||||
}
|
||||
|
||||
incr_key_base(interior_node root,
|
||||
const correction_words_array & correction_words,
|
||||
const correction_advice_array & correction_advice,
|
||||
|
|
@ -951,13 +1058,13 @@ struct incr_key_base
|
|||
correction_seeds_array correction_seeds = {})
|
||||
: leaf_nodes{std::move(leaves)},
|
||||
offset_x{offset_share},
|
||||
cmp_store_{cmp, value_cws,
|
||||
static_cast<value_cw_word>(cw_last_in),
|
||||
static_cast<value_cw_word>(cmp_addend_in),
|
||||
value_cw_coeff,
|
||||
static_cast<value_cw_word>(cw_last_coeff_in),
|
||||
tail_in, tail_coeff_in, prefix_in, prefix_coeff_in},
|
||||
public_addends{std::move(addends)},
|
||||
cmp_store_{cmp, value_cws,
|
||||
cw_word_from_u64(cw_last_in),
|
||||
cw_word_from_u64(cmp_addend_in),
|
||||
value_cw_coeff,
|
||||
cw_word_from_u64(cw_last_coeff_in),
|
||||
tail_in, tail_coeff_in, prefix_in, prefix_coeff_in},
|
||||
root_{root},
|
||||
correction_words_{correction_words},
|
||||
correction_advice_{correction_advice},
|
||||
|
|
@ -1027,6 +1134,11 @@ struct incr_key_base
|
|||
{
|
||||
cmp_store_.assign_group(delta, addend);
|
||||
}
|
||||
template <typename Mix>
|
||||
void assign_cmp_payload(Mix mix, value_cw_word addend)
|
||||
{
|
||||
cmp_store_.assign_payload(std::move(mix), addend);
|
||||
}
|
||||
HEDLEY_NO_THROW
|
||||
const prefix_cw_array & prefix_cws() const noexcept
|
||||
{
|
||||
|
|
@ -1125,9 +1237,12 @@ struct incr_key_base
|
|||
tree::traverse01_x4(parents, cw0, cw1, left, right, is_last);
|
||||
}
|
||||
|
||||
template <std::size_t I = 0>
|
||||
template <std::size_t I = 0, typename LeafT>
|
||||
HEDLEY_NO_THROW
|
||||
auto traverse_exterior(const interior_node & node) const noexcept
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_PURE
|
||||
static auto traverse_exterior(const interior_node & node,
|
||||
const LeafT & correction_word) noexcept
|
||||
{
|
||||
static_assert(num_outputs > 0, "cmp-only key has no exterior outputs");
|
||||
using Out = concrete_output_type<I>;
|
||||
|
|
@ -1154,8 +1269,50 @@ struct incr_key_base
|
|||
static_cast<psnip_uint32_t>(count),
|
||||
static_cast<psnip_uint32_t>(pos));
|
||||
}
|
||||
encode_curve_leaf_mask<Out>(mask);
|
||||
return dpf::subtract_leaf<Out>(
|
||||
dpf::get_if_lo_bit(std::get<I>(leaf_nodes).get(), node), mask);
|
||||
dpf::get_if_lo_bit(correction_word, node), mask);
|
||||
}
|
||||
|
||||
template <std::size_t I = 0>
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto traverse_exterior(const interior_node & node) const noexcept
|
||||
{
|
||||
static_assert(num_outputs > 0, "cmp-only key has no exterior outputs");
|
||||
return traverse_exterior<I>(node, std::get<I>(leaf_nodes).get());
|
||||
}
|
||||
|
||||
/// @brief Eight one-block leaves, or eight scalar leaves when the PRG is special.
|
||||
template <std::size_t I = 0, typename Out>
|
||||
void traverse_exterior_x8(const interior_node * HEDLEY_RESTRICT nodes,
|
||||
Out * HEDLEY_RESTRICT out) const noexcept
|
||||
{
|
||||
const auto & cw = std::get<I>(leaf_nodes).get();
|
||||
if constexpr (IsExtractable || meta[I].block_len != 1)
|
||||
{
|
||||
for (std::size_t t = 0; t < 8; ++t)
|
||||
out[t] = traverse_exterior<I>(nodes[t], cw);
|
||||
return;
|
||||
}
|
||||
else
|
||||
{
|
||||
using OutT = concrete_output_type<I>;
|
||||
using block = exterior_node;
|
||||
alignas(64) block seeds[8];
|
||||
alignas(64) block masks[8];
|
||||
constexpr auto pos = meta[I].pos_base
|
||||
+ meta[I].index_in_group * meta[I].block_len;
|
||||
for (std::size_t t = 0; t < 8; ++t)
|
||||
seeds[t] = utils::to_exterior_node<block>(unset_lo_2bits(nodes[t]));
|
||||
exterior_prg::eval_x8(seeds, masks, static_cast<psnip_uint32_t>(pos));
|
||||
for (std::size_t t = 0; t < 8; ++t)
|
||||
{
|
||||
encode_curve_leaf_mask<OutT>(masks[t]);
|
||||
out[t] = dpf::subtract_leaf<OutT>(
|
||||
dpf::get_if_lo_bit(cw, nodes[t]), masks[t]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
leaf_wrapper_tuple leaf_nodes;
|
||||
|
|
@ -1383,6 +1540,16 @@ using incr_dpf_key_of_t = typename incr_dpf_key_of<InteriorPRG, ExteriorPRG,
|
|||
} // namespace incr
|
||||
} // namespace detail
|
||||
|
||||
/// @brief One uniform interior node. The default root draw for distributed dealers.
|
||||
template <typename Node>
|
||||
struct uniform_node_sampler
|
||||
{
|
||||
Node operator()() const
|
||||
{
|
||||
return dpf::uniform_sample<Node>();
|
||||
}
|
||||
};
|
||||
|
||||
template <typename PRG>
|
||||
struct pseudorandom_root_sampler
|
||||
{
|
||||
|
|
@ -1390,7 +1557,10 @@ struct pseudorandom_root_sampler
|
|||
|
||||
pseudorandom_root_sampler(
|
||||
root_type && seed = dpf::uniform_sample<root_type>())
|
||||
: seed_{seed}, counter_{0} { }
|
||||
: seed_{seed}, counter_{0}
|
||||
{
|
||||
note_experiment_seed("pseudorandom_root_sampler", seed_);
|
||||
}
|
||||
|
||||
root_type operator()(psnip_uint32_t i) const
|
||||
{
|
||||
|
|
@ -1528,6 +1698,13 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
|
||||
} // namespace detail
|
||||
|
||||
/// @brief Dealer keygen. Returns the two party keys.
|
||||
/// @param args plaintext point and payloads
|
||||
/// @param root_sampler draws the interior roots
|
||||
/// @return a `party_key` pair
|
||||
/// @note Signed domains flip the MSB before the walk (`flip_msb_if_signed_integral`).
|
||||
/// @see dpf::eval_point
|
||||
/// \complexity O(n) time and O(n) key size. n is the domain bitlength (`depth`). The loop does two interior PRG expansions and writes one correction word per level, then builds one exterior leaf per output.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
|
|
|
|||
346
include/dpf/edabit.hpp
Normal file
346
include/dpf/edabit.hpp
Normal file
|
|
@ -0,0 +1,346 @@
|
|||
/// @file dpf/edabit.hpp
|
||||
/// @brief daBits and edaBits for arithmetic ↔ boolean conversion.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_EDABIT_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_EDABIT_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <stdexcept>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/ot_pack.hpp"
|
||||
#include "dpf/rss_seed.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace edabit
|
||||
{
|
||||
|
||||
/// @brief Packed boolean XOR shares + arithmetic share of `r = sum b_i 2^i`.
|
||||
template <typename Ring = std::uint64_t>
|
||||
struct edabit_share
|
||||
{
|
||||
std::vector<std::uint8_t> bits_packed; ///< ceil(ell/8) XOR / RSS-own bits
|
||||
std::vector<std::uint8_t> bits_next; ///< RSS next bit component; empty in 2PC
|
||||
Ring arith{};
|
||||
Ring arith_next{}; ///< RSS next component; 0 for 2PC
|
||||
unsigned width = 0;
|
||||
};
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
inline void set_bit(std::vector<std::uint8_t> & packed, unsigned i,
|
||||
std::uint8_t bit)
|
||||
{
|
||||
const unsigned byte = i / 8u;
|
||||
const unsigned off = i % 8u;
|
||||
if (byte >= packed.size())
|
||||
packed.resize(byte + 1u, 0);
|
||||
if (bit & 1u)
|
||||
packed[byte] = static_cast<std::uint8_t>(packed[byte] | (1u << off));
|
||||
else
|
||||
packed[byte] = static_cast<std::uint8_t>(packed[byte] & ~(1u << off));
|
||||
}
|
||||
|
||||
inline std::uint8_t get_bit(const std::vector<std::uint8_t> & packed, unsigned i)
|
||||
{
|
||||
const unsigned byte = i / 8u;
|
||||
const unsigned off = i % 8u;
|
||||
if (byte >= packed.size())
|
||||
return 0;
|
||||
return static_cast<std::uint8_t>((packed[byte] >> off) & 1u);
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
template <typename Ring = std::uint64_t>
|
||||
struct edabit_pair
|
||||
{
|
||||
edabit_share<Ring> p0;
|
||||
edabit_share<Ring> p1;
|
||||
Ring clear_r{};
|
||||
};
|
||||
|
||||
template <typename Ring = std::uint64_t>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
edabit_pair<Ring> sample_edabit_pair(unsigned ell)
|
||||
{
|
||||
if (ell == 0 || ell > 8u * sizeof(Ring))
|
||||
throw std::invalid_argument("edabit width");
|
||||
edabit_pair<Ring> out;
|
||||
out.p0.width = ell;
|
||||
out.p1.width = ell;
|
||||
const std::size_t nbytes = (ell + 7u) / 8u;
|
||||
out.p0.bits_packed.assign(nbytes, 0);
|
||||
out.p1.bits_packed.assign(nbytes, 0);
|
||||
Ring r = Ring{};
|
||||
for (unsigned i = 0; i < ell; ++i)
|
||||
{
|
||||
auto d = ot::sample_dabit_pair<Ring>();
|
||||
detail::set_bit(out.p0.bits_packed, i, d.p0.bit);
|
||||
detail::set_bit(out.p1.bits_packed, i, d.p1.bit);
|
||||
out.p0.arith = static_cast<Ring>(out.p0.arith + (d.p0.arith << i));
|
||||
out.p1.arith = static_cast<Ring>(out.p1.arith + (d.p1.arith << i));
|
||||
const Ring bit = static_cast<Ring>((d.p0.bit ^ d.p1.bit) & 1u);
|
||||
r = static_cast<Ring>(r + (bit << i));
|
||||
}
|
||||
out.clear_r = r;
|
||||
return out;
|
||||
}
|
||||
|
||||
template <typename Ring = std::uint64_t>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
edabit_share<Ring> sample_from_pack(ot::pack & pack, unsigned ell)
|
||||
{
|
||||
if (ell == 0 || ell > 8u * sizeof(Ring))
|
||||
throw std::invalid_argument("edabit width");
|
||||
edabit_share<Ring> out;
|
||||
out.width = ell;
|
||||
out.bits_packed.assign((ell + 7u) / 8u, 0);
|
||||
for (unsigned i = 0; i < ell; ++i)
|
||||
{
|
||||
auto d = pack.take_dabit<Ring>();
|
||||
detail::set_bit(out.bits_packed, i, d.bit);
|
||||
out.arith = static_cast<Ring>(out.arith + (d.arith << i));
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief Three parties' RSS edaBits: bits first, arith = sum bit_comp · 2^i.
|
||||
template <typename Ring = std::uint64_t>
|
||||
struct edabit_rss_triple
|
||||
{
|
||||
edabit_share<Ring> p0;
|
||||
edabit_share<Ring> p1;
|
||||
edabit_share<Ring> p2;
|
||||
Ring clear_r{};
|
||||
};
|
||||
|
||||
template <typename Ring = std::uint64_t>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
edabit_rss_triple<Ring> sample_rss_all(const rss::seed_bundle & bundle,
|
||||
unsigned ell, std::uint64_t index)
|
||||
{
|
||||
if (ell == 0 || ell > 8u * sizeof(Ring))
|
||||
throw std::invalid_argument("edabit width");
|
||||
edabit_rss_triple<Ring> out;
|
||||
out.p0.width = out.p1.width = out.p2.width = ell;
|
||||
const std::size_t nbytes = (ell + 7u) / 8u;
|
||||
out.p0.bits_packed.assign(nbytes, 0);
|
||||
out.p1.bits_packed.assign(nbytes, 0);
|
||||
out.p2.bits_packed.assign(nbytes, 0);
|
||||
out.p0.bits_next.assign(nbytes, 0);
|
||||
out.p1.bits_next.assign(nbytes, 0);
|
||||
out.p2.bits_next.assign(nbytes, 0);
|
||||
Ring r = Ring{};
|
||||
for (unsigned i = 0; i < ell; ++i)
|
||||
{
|
||||
// Boolean RSS: own components XOR to the bit; next matches the neighbor.
|
||||
const auto rnd = rss::random_replicated_all<std::uint8_t>(
|
||||
bundle, index + 1 + i);
|
||||
const std::uint8_t bit = static_cast<std::uint8_t>(
|
||||
(rnd.p0.own ^ rnd.p1.own ^ rnd.p2.own) & 1u);
|
||||
const std::uint8_t u = static_cast<std::uint8_t>(
|
||||
dpf::uniform_sample<std::uint8_t>() & 1u);
|
||||
const std::uint8_t v = static_cast<std::uint8_t>(
|
||||
dpf::uniform_sample<std::uint8_t>() & 1u);
|
||||
const std::uint8_t w = static_cast<std::uint8_t>(u ^ v ^ bit);
|
||||
detail::set_bit(out.p0.bits_packed, i, u);
|
||||
detail::set_bit(out.p0.bits_next, i, v);
|
||||
detail::set_bit(out.p1.bits_packed, i, v);
|
||||
detail::set_bit(out.p1.bits_next, i, w);
|
||||
detail::set_bit(out.p2.bits_packed, i, w);
|
||||
detail::set_bit(out.p2.bits_next, i, u);
|
||||
// Arithmetic RSS of the same bit: own components sum to `bit`.
|
||||
const Ring ra = dpf::uniform_sample<Ring>();
|
||||
const Ring rb = dpf::uniform_sample<Ring>();
|
||||
const Ring rc = static_cast<Ring>(Ring{bit} - ra - rb);
|
||||
out.p0.arith = static_cast<Ring>(out.p0.arith + (ra << i));
|
||||
out.p0.arith_next = static_cast<Ring>(out.p0.arith_next + (rb << i));
|
||||
out.p1.arith = static_cast<Ring>(out.p1.arith + (rb << i));
|
||||
out.p1.arith_next = static_cast<Ring>(out.p1.arith_next + (rc << i));
|
||||
out.p2.arith = static_cast<Ring>(out.p2.arith + (rc << i));
|
||||
out.p2.arith_next = static_cast<Ring>(out.p2.arith_next + (ra << i));
|
||||
r = static_cast<Ring>(r + (Ring{bit} << i));
|
||||
}
|
||||
out.clear_r = r;
|
||||
return out;
|
||||
}
|
||||
|
||||
template <typename Ring = std::uint64_t>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
edabit_share<Ring> sample_rss(const rss::seed_bundle & bundle, unsigned me,
|
||||
unsigned ell, std::uint64_t index)
|
||||
{
|
||||
auto all = sample_rss_all<Ring>(bundle, ell, index);
|
||||
if (me == 0)
|
||||
return all.p0;
|
||||
if (me == 1)
|
||||
return all.p1;
|
||||
if (me == 2)
|
||||
return all.p2;
|
||||
throw std::invalid_argument("sample_rss party");
|
||||
}
|
||||
|
||||
/// @brief Finish one GMW AND. `d_open` / `e_open` are the public `p⊕a` and `q⊕b`.
|
||||
inline std::uint8_t and_finish(const ot::bit_triple & mine, std::uint8_t d_open,
|
||||
std::uint8_t e_open, unsigned party)
|
||||
{
|
||||
std::uint8_t z = mine.c;
|
||||
z = static_cast<std::uint8_t>(z ^ (d_open & mine.b));
|
||||
z = static_cast<std::uint8_t>(z ^ (e_open & mine.a));
|
||||
if (party == 0)
|
||||
z = static_cast<std::uint8_t>(z ^ (d_open & e_open));
|
||||
return static_cast<std::uint8_t>(z & 1u);
|
||||
}
|
||||
|
||||
/// @brief Two-party AND: open the masked bits, then `and_finish` on each view.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline std::pair<std::uint8_t, std::uint8_t> and_pair(std::uint8_t p0,
|
||||
std::uint8_t p1, std::uint8_t q0, std::uint8_t q1)
|
||||
{
|
||||
auto tp = ot::sample_bit_triple_pair();
|
||||
const std::uint8_t d = static_cast<std::uint8_t>(
|
||||
(p0 ^ tp.p0.a) ^ (p1 ^ tp.p1.a));
|
||||
const std::uint8_t e = static_cast<std::uint8_t>(
|
||||
(q0 ^ tp.p0.b) ^ (q1 ^ tp.p1.b));
|
||||
return {and_finish(tp.p0, d, e, 0), and_finish(tp.p1, d, e, 1)};
|
||||
}
|
||||
|
||||
/// @brief A2B whose carry is a shared AND. Opens `x - r` and the AND masks only.
|
||||
template <typename Ring = std::uint64_t>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::pair<std::vector<std::uint8_t>, std::vector<std::uint8_t>> a2b_gmw_pair(
|
||||
const edabit_pair<Ring> & eda, Ring x0, Ring x1)
|
||||
{
|
||||
const unsigned ell = eda.p0.width;
|
||||
const Ring delta = static_cast<Ring>(
|
||||
(x0 - eda.p0.arith) + (x1 - eda.p1.arith));
|
||||
std::vector<std::uint8_t> b0((ell + 7u) / 8u, 0);
|
||||
std::vector<std::uint8_t> b1((ell + 7u) / 8u, 0);
|
||||
std::uint8_t c0 = 0;
|
||||
std::uint8_t c1 = 0;
|
||||
for (unsigned i = 0; i < ell; ++i)
|
||||
{
|
||||
const std::uint8_t r0 = detail::get_bit(eda.p0.bits_packed, i);
|
||||
const std::uint8_t r1 = detail::get_bit(eda.p1.bits_packed, i);
|
||||
const std::uint8_t di = static_cast<std::uint8_t>(
|
||||
(static_cast<std::uint64_t>(delta) >> i) & 1u);
|
||||
detail::set_bit(b0, i, static_cast<std::uint8_t>(r0 ^ di ^ c0));
|
||||
detail::set_bit(b1, i, static_cast<std::uint8_t>(r1 ^ c1));
|
||||
const std::uint8_t rd0 = di ? r0 : 0;
|
||||
const std::uint8_t rd1 = di ? r1 : 0;
|
||||
const std::uint8_t dc0 = di ? c0 : 0;
|
||||
const std::uint8_t dc1 = di ? c1 : 0;
|
||||
auto rc = and_pair(r0, r1, c0, c1);
|
||||
c0 = static_cast<std::uint8_t>(rd0 ^ dc0 ^ rc.first);
|
||||
c1 = static_cast<std::uint8_t>(rd1 ^ dc1 ^ rc.second);
|
||||
}
|
||||
return {std::move(b0), std::move(b1)};
|
||||
}
|
||||
|
||||
/// @brief Unsigned compare of XOR bit-shares. The predicate stays shared.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline std::pair<std::uint8_t, std::uint8_t> gt_bits_pair(
|
||||
const std::vector<std::uint8_t> & x0, const std::vector<std::uint8_t> & x1,
|
||||
const std::vector<std::uint8_t> & y0, const std::vector<std::uint8_t> & y1,
|
||||
unsigned n)
|
||||
{
|
||||
std::uint8_t gt0 = 0;
|
||||
std::uint8_t gt1 = 0;
|
||||
std::uint8_t eq0 = 1;
|
||||
std::uint8_t eq1 = 0;
|
||||
for (unsigned k = n; k-- > 0; )
|
||||
{
|
||||
const std::uint8_t xb0 = detail::get_bit(x0, k);
|
||||
const std::uint8_t xb1 = detail::get_bit(x1, k);
|
||||
const std::uint8_t yb0 = detail::get_bit(y0, k);
|
||||
const std::uint8_t yb1 = detail::get_bit(y1, k);
|
||||
auto xny = and_pair(xb0, xb1, static_cast<std::uint8_t>(yb0 ^ 1u), yb1);
|
||||
auto bit = and_pair(xny.first, xny.second, eq0, eq1);
|
||||
gt0 = static_cast<std::uint8_t>(gt0 ^ bit.first);
|
||||
gt1 = static_cast<std::uint8_t>(gt1 ^ bit.second);
|
||||
auto eq = and_pair(eq0, eq1, static_cast<std::uint8_t>(xb0 ^ yb0 ^ 1u),
|
||||
static_cast<std::uint8_t>(xb1 ^ yb1));
|
||||
eq0 = eq.first;
|
||||
eq1 = eq.second;
|
||||
}
|
||||
return {gt0, gt1};
|
||||
}
|
||||
|
||||
inline std::uint64_t reconstruct_bits(
|
||||
const std::vector<std::uint8_t> & a,
|
||||
const std::vector<std::uint8_t> & b, unsigned ell)
|
||||
{
|
||||
std::uint64_t v = 0;
|
||||
for (unsigned i = 0; i < ell; ++i)
|
||||
{
|
||||
const std::uint8_t bit = static_cast<std::uint8_t>(
|
||||
detail::get_bit(a, i) ^ detail::get_bit(b, i));
|
||||
v |= (static_cast<std::uint64_t>(bit) << i);
|
||||
}
|
||||
return v;
|
||||
}
|
||||
|
||||
/// @brief One party's B2A contribution after opening `mask = b ⊕ r` per bit.
|
||||
template <typename Ring = std::uint64_t>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
Ring b2a_party_bit(std::uint8_t b_share, const ot::dabit<Ring> & r,
|
||||
std::uint8_t mask_open, unsigned party, unsigned shift)
|
||||
{
|
||||
// b = r ⊕ mask. Arithmetic: r_arith + mask * (1 - 2*r_bit_as...)
|
||||
// Standard: open c = b ⊕ r; then [b] = [r] + c - 2c[r] for arith r in {0,1}.
|
||||
// With XOR r_bit matching r_arith:
|
||||
// share = r.arith + (party==0 ? Ring{mask_open} : 0)
|
||||
// - Ring{2} * Ring{mask_open} * Ring{r.bit}
|
||||
// but r.bit is XOR-shared: use r.arith which equals the bit value shares.
|
||||
Ring s = r.arith;
|
||||
if (party == 0)
|
||||
s = static_cast<Ring>(s + Ring{mask_open});
|
||||
// Subtract 2 * mask * r: each party subtracts 2*mask*r.arith (additive r).
|
||||
s = static_cast<Ring>(s - static_cast<Ring>(Ring{2} * Ring{mask_open} * r.arith));
|
||||
(void)b_share;
|
||||
return static_cast<Ring>(s << shift);
|
||||
}
|
||||
|
||||
/// @brief Two-party B2A: each uses its bit share + dabits; open mask = b⊕r.
|
||||
template <typename Ring = std::uint64_t>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::pair<Ring, Ring> b2a_pair(const std::vector<std::uint8_t> & bits0,
|
||||
const std::vector<std::uint8_t> & bits1, unsigned ell,
|
||||
ot::pack & pack0, ot::pack & pack1)
|
||||
{
|
||||
if (pack0.remaining_b2a() < ell || pack1.remaining_b2a() < ell)
|
||||
throw std::runtime_error("b2a_pair: need dabits");
|
||||
Ring s0{}, s1{};
|
||||
for (unsigned i = 0; i < ell; ++i)
|
||||
{
|
||||
auto d0 = pack0.take_dabit<Ring>();
|
||||
auto d1 = pack1.take_dabit<Ring>();
|
||||
const std::uint8_t b0 = detail::get_bit(bits0, i);
|
||||
const std::uint8_t b1 = detail::get_bit(bits1, i);
|
||||
const std::uint8_t mask = static_cast<std::uint8_t>(
|
||||
(b0 ^ d0.bit) ^ (b1 ^ d1.bit));
|
||||
s0 = static_cast<Ring>(s0 + b2a_party_bit(b0, d0, mask, 0, i));
|
||||
s1 = static_cast<Ring>(s1 + b2a_party_bit(b1, d1, mask, 1, i));
|
||||
}
|
||||
return {s0, s1};
|
||||
}
|
||||
|
||||
/// @brief Oracle expected value (not a party protocol).
|
||||
template <typename Ring = std::uint64_t>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
Ring b2a_clear(const std::vector<std::uint8_t> & bits0,
|
||||
const std::vector<std::uint8_t> & bits1, unsigned ell)
|
||||
{
|
||||
return static_cast<Ring>(reconstruct_bits(bits0, bits1, ell));
|
||||
}
|
||||
|
||||
} // namespace edabit
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_EDABIT_HPP__
|
||||
|
|
@ -2,6 +2,8 @@
|
|||
/// @brief Evaluate every input in the DPF domain.
|
||||
/// @details Equivalent to `eval_interval` from
|
||||
/// `std::numeric_limits<input_type>::min()` through `max()`.
|
||||
/// When the input offset is not yet assigned, use `defer_eval_full`
|
||||
/// (full-domain identity eval plus a deferred rotation view).
|
||||
/// @snippet evaluation/eval_full.cpp eval-full
|
||||
/// @author Ryan Henry <ryan.henry@ucalgary.ca>
|
||||
/// @author Christopher Jiang <christopher.jiang@ucalgary.ca>
|
||||
|
|
@ -27,6 +29,7 @@
|
|||
#include "dpf/interval_memoizer.hpp"
|
||||
#include "dpf/rotation_iterable.hpp"
|
||||
#include "dpf/subinterval_iterable.hpp"
|
||||
#include "dpf/verifiable.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
|
@ -41,7 +44,8 @@ template <std::size_t ...Is,
|
|||
std::size_t ...IIs,
|
||||
std::enable_if_t<dpf::is_wildcard_v<typename DpfKey::raw_input_type>, bool> = false>
|
||||
auto eval_full(const DpfKey & dpf, OutputBuffers && outbufs,
|
||||
IntervalMemoizer && memoizer, std::index_sequence<IIs...>)
|
||||
IntervalMemoizer && memoizer, std::index_sequence<IIs...>,
|
||||
proof_token * pi = nullptr)
|
||||
{
|
||||
using dpf_type = DpfKey;
|
||||
using input_type = typename dpf_type::input_type;
|
||||
|
|
@ -50,7 +54,7 @@ auto eval_full(const DpfKey & dpf, OutputBuffers && outbufs,
|
|||
dpf::internal::eval_interval_impl<Is...>(dpf,
|
||||
std::numeric_limits<input_type>::min(),
|
||||
std::numeric_limits<input_type>::max(),
|
||||
outbufs, memoizer, std::make_index_sequence<sizeof...(Is)>());
|
||||
outbufs, memoizer, std::make_index_sequence<sizeof...(Is)>(), pi);
|
||||
|
||||
return utils::make_tuple(dpf::rotation_iterable(std::begin(utils::get<IIs>(outbufs)), std::end(utils::get<IIs>(outbufs)), offset)...);
|
||||
}
|
||||
|
|
@ -62,7 +66,8 @@ template <std::size_t ...Is,
|
|||
std::size_t ...IIs,
|
||||
std::enable_if_t<!dpf::is_wildcard_v<typename DpfKey::raw_input_type>, bool> = false>
|
||||
auto eval_full(const DpfKey & dpf, OutputBuffers && outbufs,
|
||||
IntervalMemoizer && memoizer, std::index_sequence<IIs...>)
|
||||
IntervalMemoizer && memoizer, std::index_sequence<IIs...>,
|
||||
proof_token * pi = nullptr)
|
||||
{
|
||||
using dpf_type = DpfKey;
|
||||
using input_type = typename dpf_type::input_type;
|
||||
|
|
@ -70,7 +75,7 @@ auto eval_full(const DpfKey & dpf, OutputBuffers && outbufs,
|
|||
dpf::internal::eval_interval_impl<Is...>(dpf,
|
||||
std::numeric_limits<input_type>::min(),
|
||||
std::numeric_limits<input_type>::max(),
|
||||
outbufs, memoizer, std::make_index_sequence<sizeof...(Is)>());
|
||||
outbufs, memoizer, std::make_index_sequence<sizeof...(Is)>(), pi);
|
||||
|
||||
return utils::make_tuple(
|
||||
subinterval_iterable(std::begin(utils::get<IIs>(outbufs)),
|
||||
|
|
@ -83,6 +88,7 @@ auto eval_full(const DpfKey & dpf, OutputBuffers && outbufs,
|
|||
|
||||
} // namespace internal
|
||||
|
||||
/// \complexity Same expansion as `eval_interval` on the whole domain. L = 2^{n - lg(outputs_per_leaf)} leaf nodes, n = `depth`. Time Θ(L) interior traversals. The output buffer stores one slot per domain point (2^n).
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
|
|
@ -91,13 +97,60 @@ template <std::size_t I = 0,
|
|||
std::enable_if_t<looks_like_dpf_key_v<DpfKey> && !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto eval_full(const DpfKey & dpf, OutputBuffers && outbufs,
|
||||
IntervalMemoizer && memoizer)
|
||||
IntervalMemoizer && memoizer, proof_token * pi = nullptr)
|
||||
{
|
||||
assert_not_wildcard_output<I, Is...>(dpf);
|
||||
|
||||
return internal::eval_full<I, Is...>(dpf, outbufs, memoizer, std::make_index_sequence<1+sizeof...(Is)>());
|
||||
return internal::eval_full<I, Is...>(dpf, outbufs, memoizer,
|
||||
std::make_index_sequence<1+sizeof...(Is)>(), pi);
|
||||
}
|
||||
|
||||
/// @brief Evaluate the whole domain and fold a once-per-BFS-node VDPF proof.
|
||||
/// \complexity Same expansion as `eval_interval` on the whole domain. L = 2^{n - lg(outputs_per_leaf)} leaf nodes, n = `depth`. Time Θ(L) interior traversals. The output buffer stores one slot per domain point (2^n).
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
typename OutputBuffers,
|
||||
typename IntervalMemoizer,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey> && !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto eval_full(const DpfKey & dpf, OutputBuffers && outbufs,
|
||||
IntervalMemoizer && memoizer, prove_ref pr)
|
||||
{
|
||||
static_assert(DpfKey::is_verifiable,
|
||||
"eval_full(..., prove(π)): key must carry dpf::verifiable");
|
||||
detail::vdpf::init_proof(pr.token, dpf);
|
||||
auto out = eval_full<I, Is...>(dpf, std::forward<OutputBuffers>(outbufs),
|
||||
std::forward<IntervalMemoizer>(memoizer), &pr.token);
|
||||
detail::vdpf::fold_output_binding(pr.token, dpf);
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief Evaluate the whole domain and fold each written output into a sketch.
|
||||
/// \complexity Same expansion as `eval_interval` on the whole domain. L = 2^{n - lg(outputs_per_leaf)} leaf nodes, n = `depth`. Time Θ(L) interior traversals. The output buffer stores one slot per domain point (2^n).
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
typename OutputBuffers,
|
||||
typename IntervalMemoizer,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey> && !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto eval_full(const DpfKey & dpf, OutputBuffers && outbufs,
|
||||
IntervalMemoizer && memoizer, sketch_ref & sk)
|
||||
{
|
||||
static_assert(DpfKey::is_extractable,
|
||||
"eval_full(..., sketch(σ)): key must carry dpf::extractable");
|
||||
auto ret = eval_full<I, Is...>(dpf, outbufs,
|
||||
std::forward<IntervalMemoizer>(memoizer));
|
||||
if constexpr (sizeof...(Is) == 0)
|
||||
{
|
||||
for (std::size_t k = 0; k < utils::size(outbufs); ++k)
|
||||
sk.absorb(outbufs[k]);
|
||||
}
|
||||
return ret;
|
||||
}
|
||||
|
||||
/// \complexity Same expansion as `eval_interval` on the whole domain. L = 2^{n - lg(outputs_per_leaf)} leaf nodes, n = `depth`. Time Θ(L) interior traversals. The output buffer stores one slot per domain point (2^n).
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
|
|
@ -109,11 +162,33 @@ template <std::size_t I = 0,
|
|||
HEDLEY_ALWAYS_INLINE
|
||||
auto eval_full(const DpfKey & dpf, OutputBuffers & outbufs) // NOLINT(runtime/references)
|
||||
{
|
||||
using input_type = typename DpfKey::input_type;
|
||||
return eval_full<I, Is...>(dpf, outbufs,
|
||||
dpf::make_basic_full_memoizer(dpf));
|
||||
// The full-domain workspace is keyed only by the key type. Reusing it
|
||||
// drops a heap allocation per call and lets a repeated key skip the
|
||||
// interior rebuild (assign_interval keeps the last level).
|
||||
thread_local auto memo = dpf::make_basic_full_memoizer<DpfKey>();
|
||||
return eval_full<I, Is...>(dpf, outbufs, memo);
|
||||
}
|
||||
|
||||
/// @brief Evaluate the whole domain into `outbufs` and fold a VDPF proof.
|
||||
/// \complexity Same expansion as `eval_interval` on the whole domain. L = 2^{n - lg(outputs_per_leaf)} leaf nodes, n = `depth`. Time Θ(L) interior traversals. The output buffer stores one slot per domain point (2^n).
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
typename OutputBuffers,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey> && !is_multilevel_key_v<DpfKey>, bool> = true,
|
||||
std::enable_if_t<!std::is_base_of_v<
|
||||
dpf::interval_memoizer_base<unwrap_party_key_t<DpfKey>>,
|
||||
std::decay_t<OutputBuffers>>, bool> = true>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto eval_full(const DpfKey & dpf, OutputBuffers & outbufs, prove_ref pr) // NOLINT(runtime/references)
|
||||
{
|
||||
static_assert(DpfKey::is_verifiable,
|
||||
"eval_full(..., prove(π)): key must carry dpf::verifiable");
|
||||
return eval_full<I, Is...>(dpf, outbufs,
|
||||
dpf::make_basic_full_memoizer(dpf), pr);
|
||||
}
|
||||
|
||||
/// \complexity Same expansion as `eval_interval` on the whole domain. L = 2^{n - lg(outputs_per_leaf)} leaf nodes, n = `depth`. Time Θ(L) interior traversals. The output buffer stores one slot per domain point (2^n).
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
|
|
@ -144,6 +219,7 @@ auto eval_full(const DpfKey & dpf,
|
|||
/// @tparam DpfKey DPF key type
|
||||
/// @param dpf the DPF key
|
||||
/// @return the evaluation result
|
||||
/// \complexity Same expansion as `eval_interval` on the whole domain. L = 2^{n - lg(outputs_per_leaf)} leaf nodes, n = `depth`. Time Θ(L) interior traversals. The output buffer stores one slot per domain point (2^n).
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
|
|
@ -151,11 +227,47 @@ template <std::size_t I = 0,
|
|||
HEDLEY_ALWAYS_INLINE
|
||||
auto eval_full(const DpfKey & dpf)
|
||||
{
|
||||
using input_type = typename DpfKey::input_type;
|
||||
return eval_full<I, Is...>(dpf,
|
||||
dpf::make_basic_full_memoizer(dpf));
|
||||
}
|
||||
|
||||
/// @brief Full-domain deferred eval while the input offset is unset.
|
||||
/// @details Sugar for `defer_eval_interval` over `[min, max]`. `outbufs` must
|
||||
/// be full-domain sized. After assign, `.get()` yields the logical
|
||||
/// full-domain view (same values as eager `eval_full` after assign).
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
typename OutputBuffers,
|
||||
typename IntervalMemoizer,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
auto defer_eval_full(const DpfKey & dpf, OutputBuffers & outbufs,
|
||||
IntervalMemoizer && memoizer) // NOLINT(runtime/references)
|
||||
{
|
||||
using input_type = typename DpfKey::input_type;
|
||||
return defer_eval_interval<I, Is...>(dpf,
|
||||
std::numeric_limits<input_type>::min(),
|
||||
std::numeric_limits<input_type>::max(),
|
||||
outbufs, std::forward<IntervalMemoizer>(memoizer));
|
||||
}
|
||||
|
||||
/// @brief `defer_eval_full` with a basic full-domain memoizer.
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
typename OutputBuffers,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>, bool> = true,
|
||||
std::enable_if_t<!std::is_base_of_v<
|
||||
dpf::interval_memoizer_base<unwrap_party_key_t<DpfKey>>,
|
||||
std::decay_t<OutputBuffers>>, bool> = true>
|
||||
auto defer_eval_full(const DpfKey & dpf, OutputBuffers & outbufs) // NOLINT(runtime/references)
|
||||
{
|
||||
return defer_eval_full<I, Is...>(dpf, outbufs,
|
||||
dpf::make_basic_full_memoizer(dpf));
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_EVAL_FULL_HPP__
|
||||
|
|
|
|||
|
|
@ -1,16 +1,42 @@
|
|||
/// @file dpf/eval_inner_product.hpp
|
||||
/// @brief Full / interval DPF evaluation that reduces against a public
|
||||
/// weight vector instead of materializing the output.
|
||||
/// @details Same interior + batched exterior AES as `eval_interval`, but
|
||||
/// each packed leaf is multiply-accumulated into a scalar:
|
||||
/// additive outputs sum `DPF(x) * w[x]`, XOR outputs xor
|
||||
/// `DPF(x) & w[x]`. A prepared memoizer skips the interior walk
|
||||
/// so the tree can be expanded before the weights exist.
|
||||
/// @brief Full / interval / sequence DPF evaluation that reduces against a
|
||||
/// public vector instead of materializing the output.
|
||||
/// @details Three local forms share the same interval / sequence domains:
|
||||
/// - **Batched leaf walk** (no tag): one output, weights in
|
||||
/// `eval_interval` layout, exterior AES batched like that walk,
|
||||
/// O(1) accumulator. Cost shape matches `eval_interval` on the
|
||||
/// same range (Θ(L) nodes) plus a multiply-add per packed slot.
|
||||
/// - **`dpf::paired`** (row-wise): one row of the weight vector per
|
||||
/// input. A row is a scalar (one output) or a `tuple` / `array`
|
||||
/// zipped with several outputs — leaf slots or an ancestor prefix
|
||||
/// plus the leaf — read off one path. Products use `operator*`;
|
||||
/// they are summed with `operator+`. Same walk cost as the point
|
||||
/// list, plus O(1) arithmetic per selected output per input.
|
||||
/// - **`dpf::columns`** (transposed / column-wise): one output,
|
||||
/// several weight streams, one accumulator per stream, one walk.
|
||||
/// Products are *not* summed across streams. A stream is anything
|
||||
/// with `w[i]` or `w(i)`. `dpf::project` maps the share before the
|
||||
/// multiply; `dpf::also` sees the unmapped share. Cost is one path
|
||||
/// walk plus O(stream count) arithmetic per input.
|
||||
///
|
||||
/// A single-stream `columns` result matches `paired` on that stream;
|
||||
/// a single-output batched leaf walk matches `paired` when the
|
||||
/// interval is leaf-aligned (covering-leaf weights equal the clipped
|
||||
/// domain points). Unaligned intervals still weight every lane of the
|
||||
/// covering leaves — same layout as one-key `eval_interval` buffers /
|
||||
/// cohort interval inner products, not the clipped iterable. A sized
|
||||
/// weight container shorter than that covering span throws
|
||||
/// `std::invalid_argument`. `paired` and `columns` throw the same way
|
||||
/// when a sized stream is shorter than the point list or the clipped
|
||||
/// interval.
|
||||
|
||||
#ifndef LIBDPF_INCLUDE_DPF_EVAL_INNER_PRODUCT_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_EVAL_INNER_PRODUCT_HPP__
|
||||
|
||||
#include <algorithm>
|
||||
#include <array>
|
||||
#include <iterator>
|
||||
#include <vector>
|
||||
#include <cstddef>
|
||||
#include <cstring>
|
||||
#include <limits>
|
||||
|
|
@ -25,8 +51,11 @@
|
|||
|
||||
#include "dpf/dpf_key.hpp"
|
||||
#include "dpf/eval_interval.hpp"
|
||||
#include "dpf/eval_sequence.hpp"
|
||||
#include "dpf/eval_target.hpp"
|
||||
#include "dpf/incremental.hpp"
|
||||
#include "dpf/leaf_node.hpp"
|
||||
#include "dpf/path_memoizer.hpp"
|
||||
#include "dpf/twiddle.hpp"
|
||||
#include "dpf/utils.hpp"
|
||||
#include "dpf/xor_wrapper.hpp"
|
||||
|
|
@ -34,6 +63,36 @@
|
|||
namespace dpf
|
||||
{
|
||||
|
||||
namespace detail_ip_check
|
||||
{
|
||||
|
||||
template <typename T, typename = void>
|
||||
struct has_container_size : std::false_type
|
||||
{
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
struct has_container_size<T,
|
||||
std::void_t<decltype(std::size(std::declval<const T &>()))>>
|
||||
: std::true_type
|
||||
{
|
||||
};
|
||||
|
||||
/// @brief Throw when a sized weight container is shorter than the walk.
|
||||
/// Unsizable weights (`w(i)` callables, raw pointers) are left to
|
||||
/// the caller.
|
||||
template <typename W>
|
||||
void require_weight_count(const W & w, std::size_t need, const char * what)
|
||||
{
|
||||
if constexpr (has_container_size<std::decay_t<W>>::value)
|
||||
{
|
||||
if (static_cast<std::size_t>(std::size(w)) < need)
|
||||
throw std::invalid_argument(what);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace detail_ip_check
|
||||
|
||||
namespace internal
|
||||
{
|
||||
|
||||
|
|
@ -220,6 +279,25 @@ void eval_inner_product_exterior(const DpfKey & dpf, IntegralT from_node,
|
|||
auto *nodes = memoizer[DpfKey::depth];
|
||||
auto cws = std::make_tuple(std::get<Is>(dpf.leaf_nodes).get()...);
|
||||
|
||||
if constexpr (DpfKey::is_extractable)
|
||||
{
|
||||
std::size_t j = 0, k = start;
|
||||
for (; j < nodes_in_interval; ++j, ++k)
|
||||
{
|
||||
auto apply_output = [&](auto out_index, auto buf_index)
|
||||
{
|
||||
constexpr std::size_t out_i = decltype(out_index)::value;
|
||||
constexpr std::size_t buf_i = decltype(buf_index)::value;
|
||||
auto leaf = dpf.template traverse_exterior<out_i>(nodes[j]);
|
||||
std::get<buf_i>(accs).mac(leaf, k * opl, opl,
|
||||
utils::get<buf_i>(weights));
|
||||
};
|
||||
(apply_output(std::integral_constant<std::size_t, Is>{},
|
||||
std::integral_constant<std::size_t, IIs>{}), ...);
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
HEDLEY_PRAGMA(GCC diagnostic push)
|
||||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
||||
auto apply_masks = [&](std::size_t k, const node_type & node,
|
||||
|
|
@ -274,12 +352,15 @@ HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
|||
{
|
||||
alignas(node_type) node_type seeds[8];
|
||||
alignas(node_type) node_type masks[8];
|
||||
HEDLEY_PRAGMA(GCC diagnostic push)
|
||||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Warray-bounds")
|
||||
DPF_UNROLL_LOOP
|
||||
for (std::size_t t = 0; t < 8; ++t)
|
||||
{
|
||||
seeds[t] = utils::to_exterior_node<node_type>(
|
||||
unset_lo_2bits(nodes[j + t]));
|
||||
}
|
||||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||||
DpfKey::exterior_prg::eval_x8(seeds, masks, pos);
|
||||
DPF_UNROLL_LOOP
|
||||
for (std::size_t t = 0; t < 8; ++t)
|
||||
|
|
@ -374,6 +455,17 @@ auto eval_inner_product_impl(const DpfKey & dpf, InputT from, InputT to,
|
|||
utils::bitlength_of_v<InputT>);
|
||||
auto segs = utils::split_leaf_nodes(from_node, to_node, dpf.depth, wraps);
|
||||
|
||||
constexpr std::size_t opl = DpfKey::outputs_per_leaf;
|
||||
const std::size_t lanes = segs.total * opl;
|
||||
auto check_one = [&](auto which)
|
||||
{
|
||||
constexpr std::size_t wi = decltype(which)::value;
|
||||
detail_ip_check::require_weight_count(
|
||||
utils::get<wi>(weights), lanes,
|
||||
"inner product weights are shorter than the covering leaves");
|
||||
};
|
||||
(check_one(std::integral_constant<std::size_t, IIs>{}), ...);
|
||||
|
||||
auto accs = std::make_tuple(
|
||||
ip_accum<typename DpfKey::concrete_output_type<Is>>{}...);
|
||||
auto idxs = std::index_sequence<IIs...>{};
|
||||
|
|
@ -436,10 +528,11 @@ void eval_prepare_full(const DpfKey & dpf, IntervalMemoizer && memoizer)
|
|||
}
|
||||
|
||||
/// @brief `sum_x DPF_I(x) * w[x]` (additive) or `xor_x DPF_I(x) & w[x]` (XOR).
|
||||
/// @details `w[j]` is the weight for the `j`-th output in the interval, matching
|
||||
/// `eval_interval`'s destination layout. Multiple `Is` take a tuple of
|
||||
/// weight ranges and return a tuple of accumulators; a single `I` takes
|
||||
/// one range and returns one accumulator.
|
||||
/// @details `w[j]` is the weight for the `j`-th lane of the *covering* leaves
|
||||
/// of `[from, to]`, matching `eval_interval`'s destination buffer (not the
|
||||
/// clipped iterable). Multiple `Is` take a tuple of weight ranges and return
|
||||
/// a tuple of accumulators; a single `I` takes one range and returns one
|
||||
/// accumulator.
|
||||
/// @tparam I output index
|
||||
/// @tparam Is is
|
||||
/// @tparam DpfKey DPF key type
|
||||
|
|
@ -453,6 +546,7 @@ void eval_prepare_full(const DpfKey & dpf, IntervalMemoizer && memoizer)
|
|||
/// @param weights the weights
|
||||
/// @param memoizer the memoizer built for this key
|
||||
/// @return `sum_x DPF_I(x) * w[x]` (additive) or `xor_x DPF_I(x) & w[x]` (XOR)
|
||||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
|
|
@ -462,6 +556,7 @@ template <std::size_t I = 0,
|
|||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_inner_product(const DpfKey & dpf, InputT from, InputT to,
|
||||
Weights && weights, IntervalMemoizer && memoizer)
|
||||
{
|
||||
|
|
@ -471,6 +566,51 @@ auto eval_inner_product(const DpfKey & dpf, InputT from, InputT to,
|
|||
weights, memoizer, std::make_index_sequence<1 + sizeof...(Is)>{});
|
||||
}
|
||||
|
||||
/// @brief Inner product with a VDPF path proof over the same interval nodes.
|
||||
/// @details Folds once per BFS node (same transcript as `prove_interval`), then
|
||||
/// evaluates. Weights are not mixed into the token.
|
||||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
typename InputT,
|
||||
typename Weights,
|
||||
typename IntervalMemoizer,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_inner_product(const DpfKey & dpf, InputT from, InputT to,
|
||||
Weights && weights, IntervalMemoizer && memoizer, prove_ref pr)
|
||||
{
|
||||
static_assert(DpfKey::is_verifiable,
|
||||
"eval_inner_product(..., prove(π)): key must carry dpf::verifiable");
|
||||
detail::vdpf::init_proof(pr.token, dpf);
|
||||
prove_fold_interval(dpf, from, to, pr.token);
|
||||
detail::vdpf::fold_output_binding(pr.token, dpf);
|
||||
return eval_inner_product<I, Is...>(dpf, from, to,
|
||||
std::forward<Weights>(weights),
|
||||
std::forward<IntervalMemoizer>(memoizer));
|
||||
}
|
||||
|
||||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
typename InputT,
|
||||
typename Weights,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_inner_product(const DpfKey & dpf, InputT from, InputT to,
|
||||
Weights && weights, prove_ref pr)
|
||||
{
|
||||
return eval_inner_product<I, Is...>(dpf, from, to,
|
||||
std::forward<Weights>(weights),
|
||||
dpf::make_basic_interval_memoizer(dpf, from, to), pr);
|
||||
}
|
||||
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
|
|
@ -479,6 +619,7 @@ template <std::size_t I = 0,
|
|||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_full_inner_product(const DpfKey & dpf, Weights && weights,
|
||||
IntervalMemoizer && memoizer)
|
||||
{
|
||||
|
|
@ -489,6 +630,700 @@ auto eval_full_inner_product(const DpfKey & dpf, Weights && weights,
|
|||
weights, memoizer);
|
||||
}
|
||||
|
||||
/// @brief Tag: row-wise zip of several DPF outputs with each weight element.
|
||||
struct paired_t
|
||||
{
|
||||
};
|
||||
|
||||
inline constexpr paired_t paired{};
|
||||
|
||||
/// @brief Tag: transposed walk — one output, several weight streams kept apart.
|
||||
/// @details Unlike `paired`, products are not summed across streams.
|
||||
struct columns_t
|
||||
{
|
||||
};
|
||||
|
||||
inline constexpr columns_t columns{};
|
||||
|
||||
/// @brief Map a leaf share before it is multiplied by column weights.
|
||||
template <typename F>
|
||||
struct project_fn
|
||||
{
|
||||
F fn;
|
||||
};
|
||||
|
||||
/// @brief Wrap `fn` as the column projector. `fn` is called as `fn(share)`.
|
||||
template <typename F>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
project_fn<std::decay_t<F>> project(F && fn)
|
||||
{
|
||||
return project_fn<std::decay_t<F>>{std::forward<F>(fn)};
|
||||
}
|
||||
|
||||
/// @brief Observe each unmapped share during a `columns` walk.
|
||||
template <typename F>
|
||||
struct also_fn
|
||||
{
|
||||
F fn;
|
||||
};
|
||||
|
||||
/// @brief Wrap `fn` as a `columns` side visit.
|
||||
/// @details `fn` is called as `fn(i, x, share)`: list index, domain point,
|
||||
/// then the share `project` has not seen.
|
||||
template <typename F>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
also_fn<std::decay_t<F>> also(F && fn)
|
||||
{
|
||||
return also_fn<std::decay_t<F>>{std::forward<F>(fn)};
|
||||
}
|
||||
|
||||
/// @brief Column projector that returns the share unchanged.
|
||||
struct identity_project
|
||||
{
|
||||
template <typename T>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr T operator()(T value) const
|
||||
{
|
||||
return value;
|
||||
}
|
||||
};
|
||||
|
||||
/// @brief Column side visit that ignores its arguments.
|
||||
struct noop_also
|
||||
{
|
||||
template <typename... A>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void operator()(A && ...) const noexcept
|
||||
{
|
||||
}
|
||||
};
|
||||
|
||||
namespace detail_ip
|
||||
{
|
||||
|
||||
template <typename T, typename = void>
|
||||
struct has_public_addends : std::false_type
|
||||
{
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
struct has_public_addends<T, std::void_t<decltype(std::declval<const T &>().public_addends)>>
|
||||
: std::true_type
|
||||
{
|
||||
};
|
||||
|
||||
template <typename Row>
|
||||
struct is_std_array : std::false_type
|
||||
{
|
||||
};
|
||||
|
||||
template <typename T, std::size_t N>
|
||||
struct is_std_array<std::array<T, N>> : std::true_type
|
||||
{
|
||||
};
|
||||
|
||||
/// @brief Component `K` of a row. A scalar row pairs with output 0 only.
|
||||
template <std::size_t K, typename Row>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
decltype(auto) component(Row && row)
|
||||
{
|
||||
using R = std::decay_t<Row>;
|
||||
if constexpr (utils::is_tuple_v<R> || is_std_array<R>::value)
|
||||
return std::get<K>(std::forward<Row>(row));
|
||||
else
|
||||
{
|
||||
static_assert(K == 0,
|
||||
"a scalar weight pairs with one DPF output; use a tuple or "
|
||||
"std::array row for several outputs");
|
||||
return std::forward<Row>(row);
|
||||
}
|
||||
}
|
||||
|
||||
template <std::size_t I, typename Key>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto share_at(const Key & key, const typename Key::input_type & tx,
|
||||
const typename Key::interior_node & node)
|
||||
{
|
||||
constexpr auto bits = utils::bitlength_of_v<typename Key::input_type>;
|
||||
constexpr auto prefix = Key::meta[I].prefix == 0
|
||||
? bits : Key::meta[I].prefix;
|
||||
auto lane = detail::incr::lane_input(tx, prefix, bits);
|
||||
auto leaf = key.template traverse_exterior<I>(node);
|
||||
if constexpr (has_public_addends<Key>::value)
|
||||
detail::incr::absorb_public_addend_lane<I>(key, leaf, lane);
|
||||
using output_type = typename Key::template concrete_output_type<I>;
|
||||
return *make_eval_dpf_output<Key, output_type>(leaf, lane);
|
||||
}
|
||||
|
||||
template <std::size_t... Outs, typename Key, typename Tx, typename Path,
|
||||
typename Row, typename Acc, std::size_t... Ks>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void mac_row(const Key & key, const Tx & tx, Path & path, Row && row,
|
||||
Acc & acc, std::index_sequence<Ks...>)
|
||||
{
|
||||
// One exterior expansion per output per leaf/ancestor bucket. Sequential
|
||||
// and sorted queries reuse the node already on the path.
|
||||
((acc = acc + (share_at<Outs>(key, tx, path[Key::meta[Outs].tree_level])
|
||||
* component<Ks>(row))), ...);
|
||||
}
|
||||
|
||||
template <std::size_t... Outs, typename Key, typename Point, typename Rows,
|
||||
typename Acc>
|
||||
void accumulate_points(const Key & key, Point first, Point last, Rows && rows,
|
||||
Acc & acc)
|
||||
{
|
||||
constexpr std::size_t nout = sizeof...(Outs);
|
||||
constexpr std::size_t deepest = std::max({std::size_t{0},
|
||||
Key::meta[Outs].tree_level...});
|
||||
auto path = make_basic_path_memoizer<Key>();
|
||||
std::size_t i = 0;
|
||||
for (auto it = first; it != last; ++it, ++i)
|
||||
{
|
||||
auto tx = key.offset_x(*it);
|
||||
utils::flip_msb_if_signed_integral(tx);
|
||||
detail::ensure_level(key, tx, path, deepest);
|
||||
mac_row<Outs...>(key, tx, path, rows[i], acc,
|
||||
std::make_index_sequence<nout>{});
|
||||
}
|
||||
}
|
||||
|
||||
template <std::size_t I0, std::size_t... Rest>
|
||||
struct pack_first
|
||||
{
|
||||
static constexpr std::size_t value = I0;
|
||||
};
|
||||
|
||||
template <std::size_t... Outs, typename Key, typename Rows>
|
||||
auto accum_type_from(const Key &, Rows && rows)
|
||||
{
|
||||
using row_type = std::decay_t<decltype(rows[std::size_t{0}])>;
|
||||
using y0 = decltype(share_at<pack_first<Outs...>::value>(
|
||||
std::declval<const Key &>(),
|
||||
std::declval<const typename Key::input_type &>(),
|
||||
std::declval<const typename Key::interior_node &>()));
|
||||
using w0 = std::decay_t<decltype(component<0>(std::declval<row_type &>()))>;
|
||||
using acc_type = decltype(std::declval<y0>() * std::declval<w0>());
|
||||
return acc_type{};
|
||||
}
|
||||
|
||||
/// @brief `w[i]` when `w` is a range, otherwise `w(i)`.
|
||||
template <typename W>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
decltype(auto) weight_at(W && w, std::size_t i)
|
||||
{
|
||||
if constexpr (std::is_invocable_v<W &, std::size_t>)
|
||||
return w(i);
|
||||
else
|
||||
return w[i];
|
||||
}
|
||||
|
||||
template <typename Weights>
|
||||
struct is_column_pack : std::false_type
|
||||
{
|
||||
};
|
||||
|
||||
template <typename... Ts>
|
||||
struct is_column_pack<std::tuple<Ts...>> : std::true_type
|
||||
{
|
||||
};
|
||||
|
||||
template <typename T, std::size_t N>
|
||||
struct is_column_pack<std::array<T, N>> : std::true_type
|
||||
{
|
||||
};
|
||||
|
||||
template <std::size_t Out, std::size_t K, typename Key, typename Weights,
|
||||
typename Proj>
|
||||
auto column_accum_type()
|
||||
{
|
||||
using share_type = decltype(share_at<Out>(
|
||||
std::declval<const Key &>(),
|
||||
std::declval<const typename Key::input_type &>(),
|
||||
std::declval<const typename Key::interior_node &>()));
|
||||
using mapped = decltype(std::declval<Proj &>()(std::declval<share_type>()));
|
||||
using weight = decltype(weight_at(
|
||||
std::get<K>(std::declval<Weights &>()), std::size_t{0}));
|
||||
using acc_type = decltype(std::declval<mapped>() * std::declval<weight>());
|
||||
return acc_type{};
|
||||
}
|
||||
|
||||
template <std::size_t Out, typename Key, typename Point, typename Weights,
|
||||
typename Proj, typename Sink, std::size_t... Ks>
|
||||
auto accumulate_columns(const Key & key, Point first, Point last,
|
||||
Weights && weights, Proj && proj, Sink && sink,
|
||||
std::index_sequence<Ks...>)
|
||||
{
|
||||
static_assert(sizeof...(Ks) > 0, "columns needs at least one weight stream");
|
||||
using acc_tuple = std::tuple<decltype(column_accum_type<Out, Ks, Key,
|
||||
std::decay_t<Weights>, std::decay_t<Proj>>())...>;
|
||||
acc_tuple acc{};
|
||||
if constexpr (std::is_base_of_v<std::random_access_iterator_tag,
|
||||
typename std::iterator_traits<Point>::iterator_category>)
|
||||
{
|
||||
const auto n = static_cast<std::size_t>(std::distance(first, last));
|
||||
auto guard = [&](auto which)
|
||||
{
|
||||
constexpr std::size_t k = decltype(which)::value;
|
||||
detail_ip_check::require_weight_count(std::get<k>(weights), n,
|
||||
"column weights are shorter than the point list");
|
||||
};
|
||||
(guard(std::integral_constant<std::size_t, Ks>{}), ...);
|
||||
}
|
||||
constexpr std::size_t deepest = Key::meta[Out].tree_level;
|
||||
auto path = make_basic_path_memoizer<Key>();
|
||||
std::size_t i = 0;
|
||||
for (auto it = first; it != last; ++it, ++i)
|
||||
{
|
||||
auto tx = key.offset_x(*it);
|
||||
utils::flip_msb_if_signed_integral(tx);
|
||||
detail::ensure_level(key, tx, path, deepest);
|
||||
auto share = share_at<Out>(key, tx, path[Key::meta[Out].tree_level]);
|
||||
sink(i, *it, share);
|
||||
auto y = proj(share);
|
||||
((std::get<Ks>(acc) = std::get<Ks>(acc)
|
||||
+ (y * weight_at(std::get<Ks>(weights), i))), ...);
|
||||
}
|
||||
return acc;
|
||||
}
|
||||
|
||||
} // namespace detail_ip
|
||||
|
||||
namespace internal_paired
|
||||
{
|
||||
|
||||
template <typename Input, typename Fn>
|
||||
void for_inclusive(Input from, Input to, Fn && fn)
|
||||
{
|
||||
constexpr auto to_int = utils::to_integral_type<Input>{};
|
||||
using integral = decltype(to_int(from));
|
||||
const bool wraps = utils::interval_wraps(
|
||||
static_cast<integral>(to_int(from)),
|
||||
static_cast<integral>(to_int(to)),
|
||||
utils::bitlength_of_v<Input>);
|
||||
auto step = [&](Input x) { fn(x); };
|
||||
if (!wraps)
|
||||
{
|
||||
for (auto x = from;; ++x)
|
||||
{
|
||||
step(x);
|
||||
if (x == to)
|
||||
break;
|
||||
}
|
||||
return;
|
||||
}
|
||||
const auto hi = std::numeric_limits<Input>::max();
|
||||
const auto lo = std::numeric_limits<Input>::min();
|
||||
for (auto x = from;; ++x)
|
||||
{
|
||||
step(x);
|
||||
if (x == hi)
|
||||
break;
|
||||
}
|
||||
for (auto x = lo;; ++x)
|
||||
{
|
||||
step(x);
|
||||
if (x == to)
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
template <std::size_t... Outs, typename Key, typename Rows>
|
||||
auto run_points(const Key & key, const std::vector<typename Key::input_type> & xs,
|
||||
Rows && rows)
|
||||
{
|
||||
assert_not_wildcard_output<Outs...>(key);
|
||||
if (xs.size() == 0)
|
||||
{
|
||||
using acc_type = decltype(detail_ip::accum_type_from<Outs...>(key, rows));
|
||||
return acc_type{};
|
||||
}
|
||||
detail_ip_check::require_weight_count(rows, xs.size(),
|
||||
"paired weights are shorter than the point list");
|
||||
using acc_type = decltype(detail_ip::accum_type_from<Outs...>(key, rows));
|
||||
acc_type acc{};
|
||||
detail_ip::accumulate_points<Outs...>(key, xs.begin(), xs.end(), rows, acc);
|
||||
return acc;
|
||||
}
|
||||
|
||||
/// @brief Clipped interval dot via the interval tree, not one path per point.
|
||||
/// Wrapping intervals stay on the path walk.
|
||||
template <std::size_t I, typename DpfKey, typename InputT, typename Rows>
|
||||
auto interval_scalar(const DpfKey & dpf, InputT from, InputT to, Rows && rows)
|
||||
{
|
||||
constexpr auto to_int = utils::to_integral_type<InputT>{};
|
||||
using integral = decltype(to_int(from));
|
||||
const bool wraps = utils::interval_wraps(
|
||||
static_cast<integral>(to_int(from)),
|
||||
static_cast<integral>(to_int(to)),
|
||||
utils::bitlength_of_v<InputT>);
|
||||
if (wraps)
|
||||
{
|
||||
std::vector<InputT> xs;
|
||||
for_inclusive(from, to, [&](InputT x) { xs.push_back(x); });
|
||||
return run_points<I>(dpf, xs, std::forward<Rows>(rows));
|
||||
}
|
||||
const auto npoints = static_cast<std::size_t>(
|
||||
static_cast<integral>(to_int(to)) - static_cast<integral>(to_int(from)))
|
||||
+ std::size_t{1};
|
||||
detail_ip_check::require_weight_count(rows, npoints,
|
||||
"paired weights are shorter than the interval");
|
||||
auto buf = make_output_buffer_for_interval<I>(dpf, from, to);
|
||||
auto iter = eval_interval<I>(dpf, from, to, buf);
|
||||
using share_t = std::decay_t<decltype(*iter.begin())>;
|
||||
using weight_t = std::decay_t<decltype(rows[std::size_t{0}])>;
|
||||
using acc_t = decltype(std::declval<share_t>() * std::declval<weight_t>());
|
||||
acc_t acc{};
|
||||
std::size_t i = 0;
|
||||
for (auto it = iter.begin(); it != iter.end(); ++it, ++i)
|
||||
acc = acc + (*it * rows[i]);
|
||||
return acc;
|
||||
}
|
||||
|
||||
} // namespace internal_paired
|
||||
|
||||
/// @brief `sum_i Σ_k DPF_{out_k}(x_i) * row_i[k]` over `[from, to]`.
|
||||
/// @details One row of `rows` per input, in interval order (wrapping the same
|
||||
/// way as `eval_interval`). A row is a scalar when one output is
|
||||
/// selected, or a `std::tuple` / `std::array` with one component per
|
||||
/// output. Ancestor slots and leaf slots are read off the same path.
|
||||
/// Differs from the batched leaf walk: one path step per domain point
|
||||
/// (reuse via the path memoizer), not batched exterior AES over leaf
|
||||
/// nodes. For a single output the opened result matches the batched
|
||||
/// form when `rows` is the interval weight vector.
|
||||
/// \complexity One path ensure per input up to the deepest selected output,
|
||||
/// plus one exterior expand and multiply-add per selected output per
|
||||
/// input. Accumulator is O(1); no output buffer.
|
||||
template <std::size_t I = 0, std::size_t... Is, typename DpfKey, typename InputT,
|
||||
typename Rows>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_inner_product(paired_t, const DpfKey & dpf, InputT from, InputT to,
|
||||
Rows && rows)
|
||||
{
|
||||
using row_t = std::decay_t<decltype(std::declval<Rows &>()[std::size_t{0}])>;
|
||||
constexpr bool scalar = !utils::is_tuple_v<row_t>
|
||||
&& !detail_ip::is_std_array<row_t>::value;
|
||||
if constexpr (scalar && sizeof...(Is) == 0 && !is_multilevel_key_v<DpfKey>)
|
||||
{
|
||||
return internal_paired::interval_scalar<I>(dpf, from, to,
|
||||
std::forward<Rows>(rows));
|
||||
}
|
||||
else
|
||||
{
|
||||
std::vector<InputT> xs;
|
||||
internal_paired::for_inclusive(from, to, [&](InputT x) { xs.push_back(x); });
|
||||
return internal_paired::run_points<I, Is...>(dpf, xs,
|
||||
std::forward<Rows>(rows));
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Paired inner product over the whole input domain.
|
||||
template <std::size_t I = 0, std::size_t... Is, typename DpfKey, typename Rows>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_full_inner_product(paired_t, const DpfKey & dpf, Rows && rows)
|
||||
{
|
||||
using input_type = typename DpfKey::input_type;
|
||||
return eval_inner_product<I, Is...>(paired, dpf,
|
||||
std::numeric_limits<input_type>::min(),
|
||||
std::numeric_limits<input_type>::max(),
|
||||
std::forward<Rows>(rows));
|
||||
}
|
||||
|
||||
/// @brief Paired inner product over a sorted point list.
|
||||
/// @details Same order and sortedness rule as `eval_sequence`. Each point
|
||||
/// pairs with `rows[i]`.
|
||||
template <std::size_t I = 0, std::size_t... Is, typename DpfKey,
|
||||
typename ForwardIterator, typename Rows>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_sequence_inner_product(const DpfKey & dpf, ForwardIterator begin,
|
||||
ForwardIterator end, Rows && rows)
|
||||
{
|
||||
if (!std::is_sorted(begin, end))
|
||||
throw std::runtime_error("list must be sorted");
|
||||
std::vector<typename DpfKey::input_type> xs(begin, end);
|
||||
return internal_paired::run_points<I, Is...>(dpf, xs,
|
||||
std::forward<Rows>(rows));
|
||||
}
|
||||
|
||||
/// @brief Paired inner product over a `sequence_recipe`'s points.
|
||||
/// @details `points` is the same sorted list the recipe was built from. The
|
||||
/// recipe drives nothing the path walk does not already share; it
|
||||
/// checks that the list still matches the recipe's output count.
|
||||
template <std::size_t I = 0, std::size_t... Is, typename DpfKey,
|
||||
typename ForwardIterator, typename Rows>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_sequence_inner_product(const DpfKey & dpf,
|
||||
const sequence_recipe & recipe, ForwardIterator begin, ForwardIterator end,
|
||||
Rows && rows)
|
||||
{
|
||||
const auto n = static_cast<std::size_t>(std::distance(begin, end));
|
||||
if (n != recipe.output_indices().size())
|
||||
throw std::invalid_argument(
|
||||
"eval_sequence_inner_product: recipe and point list differ");
|
||||
return eval_sequence_inner_product<I, Is...>(dpf, begin, end,
|
||||
std::forward<Rows>(rows));
|
||||
}
|
||||
|
||||
namespace detail_columns
|
||||
{
|
||||
|
||||
template <std::size_t I, typename DpfKey, typename ForwardIterator,
|
||||
typename Weights, typename Proj, typename Sink>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto run(const DpfKey & dpf, ForwardIterator begin, ForwardIterator end,
|
||||
Weights && weights, Proj && proj, Sink && sink, bool require_sorted)
|
||||
{
|
||||
using pack = std::decay_t<Weights>;
|
||||
static_assert(detail_ip::is_column_pack<pack>::value,
|
||||
"columns weights are a std::tuple or std::array of streams; "
|
||||
"each stream is w[i] or w(i)");
|
||||
if (require_sorted && !std::is_sorted(begin, end))
|
||||
throw std::runtime_error("list must be sorted");
|
||||
std::vector<typename DpfKey::input_type> xs(begin, end);
|
||||
constexpr auto n = std::tuple_size<pack>::value;
|
||||
return detail_ip::accumulate_columns<I>(dpf, xs.begin(), xs.end(),
|
||||
std::forward<Weights>(weights), std::forward<Proj>(proj),
|
||||
std::forward<Sink>(sink), std::make_index_sequence<n>{});
|
||||
}
|
||||
|
||||
template <std::size_t I, typename DpfKey, typename InputT, typename Weights,
|
||||
std::size_t... Ks>
|
||||
auto interval_dots(const DpfKey & dpf, InputT from, InputT to, Weights && weights,
|
||||
std::index_sequence<Ks...>)
|
||||
{
|
||||
constexpr auto to_int = utils::to_integral_type<InputT>{};
|
||||
using integral = decltype(to_int(from));
|
||||
const bool wraps = utils::interval_wraps(
|
||||
static_cast<integral>(to_int(from)),
|
||||
static_cast<integral>(to_int(to)),
|
||||
utils::bitlength_of_v<InputT>);
|
||||
if (wraps)
|
||||
{
|
||||
std::vector<InputT> xs;
|
||||
internal_paired::for_inclusive(from, to, [&](InputT x) { xs.push_back(x); });
|
||||
return detail_ip::accumulate_columns<I>(dpf, xs.begin(), xs.end(),
|
||||
std::forward<Weights>(weights), identity_project{},
|
||||
noop_also{}, std::index_sequence<Ks...>{});
|
||||
}
|
||||
const auto npoints = static_cast<std::size_t>(
|
||||
static_cast<integral>(to_int(to)) - static_cast<integral>(to_int(from)))
|
||||
+ std::size_t{1};
|
||||
auto guard = [&](auto which)
|
||||
{
|
||||
constexpr std::size_t k = decltype(which)::value;
|
||||
detail_ip_check::require_weight_count(std::get<k>(weights), npoints,
|
||||
"column weights are shorter than the interval");
|
||||
};
|
||||
(guard(std::integral_constant<std::size_t, Ks>{}), ...);
|
||||
auto buf = make_output_buffer_for_interval<I>(dpf, from, to);
|
||||
auto iter = eval_interval<I>(dpf, from, to, buf);
|
||||
using acc_tuple = std::tuple<decltype(detail_ip::column_accum_type<I, Ks,
|
||||
DpfKey, std::decay_t<Weights>, identity_project>())...>;
|
||||
acc_tuple acc{};
|
||||
std::size_t i = 0;
|
||||
for (auto it = iter.begin(); it != iter.end(); ++it, ++i)
|
||||
{
|
||||
auto y = *it;
|
||||
((std::get<Ks>(acc) = std::get<Ks>(acc)
|
||||
+ (y * detail_ip::weight_at(std::get<Ks>(weights), i))), ...);
|
||||
}
|
||||
return acc;
|
||||
}
|
||||
|
||||
} // namespace detail_columns
|
||||
|
||||
/// @brief Several independent dots of one output, one walk (transposed form).
|
||||
/// @details `weights` is a `std::tuple` or `std::array` of streams. Stream
|
||||
/// `k` is either `w[i]` or `w(i)`, `i` the position in the point
|
||||
/// list (or in the interval, wrapping the same way as
|
||||
/// `eval_interval`). The result is a tuple of accumulators,
|
||||
/// `acc_k = sum_i project(DPF(x_i)) * stream_k(i)`.
|
||||
/// `dpf::project(fn)` maps the share first. `dpf::also(fn)` is
|
||||
/// called as `fn(i, x, share)` on the unmapped share. Pass either
|
||||
/// tag, both, or neither. `also` then `project` is accepted too.
|
||||
/// Sequence points are sorted nondecreasing, same as `eval_sequence`.
|
||||
/// Relative to `paired`: same path walk for one output, but streams
|
||||
/// stay separate (no cross-stream sum). One stream equals `paired`
|
||||
/// on that stream. Relative to the batched leaf walk: path-per-point
|
||||
/// instead of batched exterior AES; same opened scalar when the
|
||||
/// single stream matches the interval weight layout.
|
||||
/// \complexity One path ensure and one exterior expand per input, plus
|
||||
/// O(stream count) multiply-adds per input. Accumulators are O(stream count).
|
||||
/// @{
|
||||
|
||||
template <std::size_t I = 0, typename DpfKey, typename ForwardIterator,
|
||||
typename Weights>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_sequence_inner_product(columns_t, const DpfKey & dpf,
|
||||
ForwardIterator begin, ForwardIterator end, Weights && weights)
|
||||
{
|
||||
assert_not_wildcard_output<I>(dpf);
|
||||
return detail_columns::run<I>(dpf, begin, end,
|
||||
std::forward<Weights>(weights), identity_project{}, noop_also{}, true);
|
||||
}
|
||||
|
||||
template <std::size_t I = 0, typename DpfKey, typename ForwardIterator,
|
||||
typename Weights, typename Proj>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_sequence_inner_product(columns_t, const DpfKey & dpf,
|
||||
ForwardIterator begin, ForwardIterator end, Weights && weights,
|
||||
project_fn<Proj> proj)
|
||||
{
|
||||
assert_not_wildcard_output<I>(dpf);
|
||||
return detail_columns::run<I>(dpf, begin, end,
|
||||
std::forward<Weights>(weights), proj.fn, noop_also{}, true);
|
||||
}
|
||||
|
||||
template <std::size_t I = 0, typename DpfKey, typename ForwardIterator,
|
||||
typename Weights, typename Sink>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_sequence_inner_product(columns_t, const DpfKey & dpf,
|
||||
ForwardIterator begin, ForwardIterator end, Weights && weights,
|
||||
also_fn<Sink> sink)
|
||||
{
|
||||
assert_not_wildcard_output<I>(dpf);
|
||||
return detail_columns::run<I>(dpf, begin, end,
|
||||
std::forward<Weights>(weights), identity_project{}, sink.fn, true);
|
||||
}
|
||||
|
||||
template <std::size_t I = 0, typename DpfKey, typename ForwardIterator,
|
||||
typename Weights, typename Proj, typename Sink>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_sequence_inner_product(columns_t, const DpfKey & dpf,
|
||||
ForwardIterator begin, ForwardIterator end, Weights && weights,
|
||||
project_fn<Proj> proj, also_fn<Sink> sink)
|
||||
{
|
||||
assert_not_wildcard_output<I>(dpf);
|
||||
return detail_columns::run<I>(dpf, begin, end,
|
||||
std::forward<Weights>(weights), proj.fn, sink.fn, true);
|
||||
}
|
||||
|
||||
template <std::size_t I = 0, typename DpfKey, typename ForwardIterator,
|
||||
typename Weights, typename Proj, typename Sink>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_sequence_inner_product(columns_t, const DpfKey & dpf,
|
||||
ForwardIterator begin, ForwardIterator end, Weights && weights,
|
||||
also_fn<Sink> sink, project_fn<Proj> proj)
|
||||
{
|
||||
return eval_sequence_inner_product<I>(columns, dpf, begin, end,
|
||||
std::forward<Weights>(weights), std::move(proj), std::move(sink));
|
||||
}
|
||||
|
||||
template <std::size_t I = 0, typename DpfKey, typename ForwardIterator,
|
||||
typename Weights, typename... Extra>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_sequence_inner_product(columns_t, const DpfKey & dpf,
|
||||
const sequence_recipe & recipe, ForwardIterator begin,
|
||||
ForwardIterator end, Weights && weights, Extra && ... extra)
|
||||
{
|
||||
const auto n = static_cast<std::size_t>(std::distance(begin, end));
|
||||
if (n != recipe.output_indices().size())
|
||||
throw std::invalid_argument(
|
||||
"eval_sequence_inner_product: recipe and point list differ");
|
||||
return eval_sequence_inner_product<I>(columns, dpf, begin, end,
|
||||
std::forward<Weights>(weights), std::forward<Extra>(extra)...);
|
||||
}
|
||||
|
||||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||||
template <std::size_t I = 0, typename DpfKey, typename InputT, typename Weights,
|
||||
typename Proj, typename Sink>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_inner_product(columns_t, const DpfKey & dpf, InputT from, InputT to,
|
||||
Weights && weights, Proj && proj, Sink && sink)
|
||||
{
|
||||
assert_not_wildcard_output<I>(dpf);
|
||||
if constexpr (std::is_same_v<std::decay_t<Proj>, identity_project>
|
||||
&& std::is_same_v<std::decay_t<Sink>, noop_also>
|
||||
&& !is_multilevel_key_v<DpfKey>)
|
||||
{
|
||||
using pack = std::decay_t<Weights>;
|
||||
static_assert(detail_ip::is_column_pack<pack>::value,
|
||||
"columns weights are a std::tuple or std::array of streams; "
|
||||
"each stream is w[i] or w(i)");
|
||||
constexpr auto nstreams = std::tuple_size<pack>::value;
|
||||
return detail_columns::interval_dots<I>(dpf, from, to,
|
||||
std::forward<Weights>(weights), std::make_index_sequence<nstreams>{});
|
||||
}
|
||||
else
|
||||
{
|
||||
std::vector<InputT> xs;
|
||||
internal_paired::for_inclusive(from, to, [&](InputT x) { xs.push_back(x); });
|
||||
return detail_columns::run<I>(dpf, xs.begin(), xs.end(),
|
||||
std::forward<Weights>(weights), std::forward<Proj>(proj),
|
||||
std::forward<Sink>(sink), false);
|
||||
}
|
||||
}
|
||||
|
||||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||||
template <std::size_t I = 0, typename DpfKey, typename InputT, typename Weights>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_inner_product(columns_t, const DpfKey & dpf, InputT from, InputT to,
|
||||
Weights && weights)
|
||||
{
|
||||
return eval_inner_product<I>(columns, dpf, from, to,
|
||||
std::forward<Weights>(weights), identity_project{}, noop_also{});
|
||||
}
|
||||
|
||||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||||
template <std::size_t I = 0, typename DpfKey, typename InputT, typename Weights,
|
||||
typename Proj>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_inner_product(columns_t, const DpfKey & dpf, InputT from, InputT to,
|
||||
Weights && weights, project_fn<Proj> proj)
|
||||
{
|
||||
return eval_inner_product<I>(columns, dpf, from, to,
|
||||
std::forward<Weights>(weights), proj.fn, noop_also{});
|
||||
}
|
||||
|
||||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||||
template <std::size_t I = 0, typename DpfKey, typename InputT, typename Weights,
|
||||
typename Sink>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_inner_product(columns_t, const DpfKey & dpf, InputT from, InputT to,
|
||||
Weights && weights, also_fn<Sink> sink)
|
||||
{
|
||||
return eval_inner_product<I>(columns, dpf, from, to,
|
||||
std::forward<Weights>(weights), identity_project{}, sink.fn);
|
||||
}
|
||||
|
||||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||||
template <std::size_t I = 0, typename DpfKey, typename InputT, typename Weights,
|
||||
typename Proj, typename Sink>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_inner_product(columns_t, const DpfKey & dpf, InputT from, InputT to,
|
||||
Weights && weights, project_fn<Proj> proj, also_fn<Sink> sink)
|
||||
{
|
||||
return eval_inner_product<I>(columns, dpf, from, to,
|
||||
std::forward<Weights>(weights), proj.fn, sink.fn);
|
||||
}
|
||||
|
||||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||||
template <std::size_t I = 0, typename DpfKey, typename InputT, typename Weights,
|
||||
typename Proj, typename Sink>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_inner_product(columns_t, const DpfKey & dpf, InputT from, InputT to,
|
||||
Weights && weights, also_fn<Sink> sink, project_fn<Proj> proj)
|
||||
{
|
||||
return eval_inner_product<I>(columns, dpf, from, to,
|
||||
std::forward<Weights>(weights), proj.fn, sink.fn);
|
||||
}
|
||||
|
||||
/// @brief `columns` over the whole input domain. Same streams as the interval form.
|
||||
template <std::size_t I = 0, typename DpfKey, typename Weights, typename... Extra>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_full_inner_product(columns_t, const DpfKey & dpf, Weights && weights,
|
||||
Extra && ... extra)
|
||||
{
|
||||
using input_type = typename DpfKey::input_type;
|
||||
return eval_inner_product<I>(columns, dpf,
|
||||
std::numeric_limits<input_type>::min(),
|
||||
std::numeric_limits<input_type>::max(),
|
||||
std::forward<Weights>(weights), std::forward<Extra>(extra)...);
|
||||
}
|
||||
|
||||
/// @}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_EVAL_INNER_PRODUCT_HPP__
|
||||
|
|
|
|||
|
|
@ -4,6 +4,14 @@
|
|||
/// share per input, in that order. Pass a named output buffer;
|
||||
/// this overload binds it as a non-const reference. An interval
|
||||
/// memoizer is optional and comes after the buffer.
|
||||
///
|
||||
/// Eager evaluation requires an assigned input offset: the range is
|
||||
/// traversed at `offset_x(from)..offset_x(to)`. When the input is
|
||||
/// still a wildcard, call `defer_eval_interval` instead — that fills
|
||||
/// a **full-domain** buffer at identity and returns a
|
||||
/// `deferred_rotated_subinterval` that applies the rotation after
|
||||
/// `assign_wildcard_input`. Interior-only prep with an assigned
|
||||
/// input but unassigned leaf is `defer_traverse_interval`.
|
||||
/// @snippet evaluation/eval_interval.cpp eval-interval
|
||||
/// @author Ryan Henry <ryan.henry@ucalgary.ca>
|
||||
/// @author Christopher Jiang <christopher.jiang@ucalgary.ca>
|
||||
|
|
@ -20,6 +28,7 @@
|
|||
|
||||
#include <cstddef>
|
||||
#include <cstring>
|
||||
#include <limits>
|
||||
#include <stdexcept>
|
||||
#include <array>
|
||||
#include <tuple>
|
||||
|
|
@ -33,6 +42,9 @@
|
|||
#include "dpf/output_buffer.hpp"
|
||||
#include "dpf/interval_memoizer.hpp"
|
||||
#include "dpf/subinterval_iterable.hpp"
|
||||
#include "dpf/deferred_rotated_subinterval.hpp"
|
||||
#include "dpf/verifiable.hpp"
|
||||
#include "dpf/wildcard.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
|
@ -40,12 +52,35 @@ namespace dpf
|
|||
namespace internal
|
||||
{
|
||||
|
||||
/// @brief Fold every node at `level_index` of a truncated interval tree into `pi`.
|
||||
/// @details Once-per-BFS-node absorption: node `i` has prefix
|
||||
/// `(from_node >> (depth - level_index)) + i`. Matches the contiguous
|
||||
/// layout built by `eval_interval_interior`.
|
||||
template <typename DpfKey, typename IntegralT, typename NodeT>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void fold_interval_level(proof_token & pi, const DpfKey & dpf,
|
||||
std::size_t level_index, IntegralT from_node, std::size_t nodes_at_level,
|
||||
const NodeT * curr)
|
||||
{
|
||||
if constexpr (!DpfKey::is_verifiable)
|
||||
return;
|
||||
if (level_index == 0 || nodes_at_level == 0)
|
||||
return;
|
||||
const auto start = static_cast<psnip_uint64_t>(
|
||||
utils::shift_right(from_node, DpfKey::depth - level_index));
|
||||
const auto & cs = dpf.correction_seeds()[level_index - 1];
|
||||
for (std::size_t i = 0; i < nodes_at_level; ++i)
|
||||
{
|
||||
detail::vdpf::fold_node(pi, level_index - 1, start + i, curr[i], cs);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename DpfKey,
|
||||
typename IntervalMemoizer,
|
||||
typename IntegralT = typename DpfKey::integral_type>
|
||||
inline auto eval_interval_interior(const DpfKey & dpf, IntegralT from_node,
|
||||
IntegralT to_node, IntervalMemoizer & memoizer, // NOLINT(runtime/references)
|
||||
std::size_t to_level = DpfKey::depth)
|
||||
std::size_t to_level = DpfKey::depth, proof_token * pi = nullptr)
|
||||
{
|
||||
using dpf_type = DpfKey;
|
||||
using integral_type = typename DpfKey::integral_type;
|
||||
|
|
@ -54,6 +89,11 @@ inline auto eval_interval_interior(const DpfKey & dpf, IntegralT from_node,
|
|||
// level_index represents the current level being built
|
||||
// level_index = 0 => root
|
||||
// level_index = depth => last layer of interior nodes
|
||||
// Proving needs every truncated-tree node: a warm memoizer that resumes
|
||||
// past level 1 would skip upper folds (and basic memoizers discard them).
|
||||
if (pi != nullptr)
|
||||
memoizer.clear_assignment();
|
||||
|
||||
std::size_t level_index = memoizer.assign_interval(dpf, from_node, to_node);
|
||||
std::size_t nodes_at_level = memoizer.get_nodes_at_level();
|
||||
integral_type mask = utils::get_node_mask<dpf_type>(dpf.msb_mask, level_index);
|
||||
|
|
@ -115,6 +155,12 @@ inline auto eval_interval_interior(const DpfKey & dpf, IntegralT from_node,
|
|||
{
|
||||
curr[i] = dpf_type::traverse_interior(prev[j], cw[0], 0, is_last);
|
||||
}
|
||||
|
||||
if (pi != nullptr)
|
||||
{
|
||||
fold_interval_level(*pi, dpf, level_index, from_node,
|
||||
nodes_at_level, curr);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -136,23 +182,38 @@ inline auto eval_interval_exterior(const DpfKey & dpf, IntegralT from_node,
|
|||
|
||||
std::size_t nodes_in_interval = static_cast<std::size_t>(to_node - from_node);
|
||||
|
||||
HEDLEY_PRAGMA(GCC diagnostic push)
|
||||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
||||
auto cw = std::get<I>(dpf.leaf_nodes).get();
|
||||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||||
auto *nodes = memoizer[dpf_type::depth];
|
||||
DPF_UNROLL_LOOP
|
||||
for (std::size_t j = 0, k = start; j < nodes_in_interval; ++j, ++k)
|
||||
std::size_t j = 0, k = start;
|
||||
constexpr bool batch_leaves =
|
||||
dpf::block_length_of_leaf_v<output_type, typename DpfKey::interior_node> == 1
|
||||
&& !utils::is_packed_subbyte_v<output_type>
|
||||
&& dpf_type::outputs_per_leaf == 1;
|
||||
if constexpr (batch_leaves)
|
||||
{
|
||||
auto leaf = dpf.template traverse_exterior<I>(nodes[j],
|
||||
get_if_lo_bit(cw, nodes[j]));
|
||||
using leaf_ret = decltype(dpf.template traverse_exterior<I>(nodes[0]));
|
||||
while (j + 8 <= nodes_in_interval)
|
||||
{
|
||||
leaf_ret leaves[8];
|
||||
dpf.template traverse_exterior_x8<I>(nodes + j, leaves);
|
||||
for (std::size_t t = 0; t < 8; ++t, ++k)
|
||||
{
|
||||
utils::raw_memcpy(&outbuf[k], &leaves[t], sizeof(output_type));
|
||||
}
|
||||
j += 8;
|
||||
}
|
||||
}
|
||||
for (; j < nodes_in_interval; ++j, ++k)
|
||||
{
|
||||
// 1-arg member works for classic and verifiable/incr keys; the static
|
||||
// 2-arg form is classic-only.
|
||||
auto leaf = dpf.template traverse_exterior<I>(nodes[j]);
|
||||
if constexpr (utils::is_packed_subbyte_v<output_type>)
|
||||
{
|
||||
store_leaf_bytes(outbuf, k, leaf);
|
||||
}
|
||||
else
|
||||
{
|
||||
std::memcpy(&outbuf[k*dpf_type::outputs_per_leaf], &leaf,
|
||||
utils::raw_memcpy(&outbuf[k*dpf_type::outputs_per_leaf], &leaf,
|
||||
sizeof(output_type) * dpf_type::outputs_per_leaf);
|
||||
}
|
||||
}
|
||||
|
|
@ -174,7 +235,7 @@ void store_interval_leaf(OutputBuffer && outbuf, std::size_t k, const LeafT & le
|
|||
}
|
||||
else
|
||||
{
|
||||
std::memcpy(&outbuf[k * dpf_type::outputs_per_leaf], &leaf,
|
||||
utils::raw_memcpy(&outbuf[k * dpf_type::outputs_per_leaf], &leaf,
|
||||
sizeof(output_type) * dpf_type::outputs_per_leaf);
|
||||
}
|
||||
}
|
||||
|
|
@ -273,12 +334,15 @@ HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
|||
{
|
||||
alignas(node_type) node_type seeds[8];
|
||||
alignas(node_type) node_type masks[8];
|
||||
HEDLEY_PRAGMA(GCC diagnostic push)
|
||||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Warray-bounds")
|
||||
DPF_UNROLL_LOOP
|
||||
for (std::size_t t = 0; t < 8; ++t)
|
||||
{
|
||||
seeds[t] = utils::to_exterior_node<node_type>(
|
||||
unset_lo_2bits(nodes[j + t]));
|
||||
}
|
||||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||||
DpfKey::exterior_prg::eval_x8(seeds, masks, pos);
|
||||
DPF_UNROLL_LOOP
|
||||
for (std::size_t t = 0; t < 8; ++t)
|
||||
|
|
@ -330,22 +394,28 @@ void eval_interval_exterior_all(const DpfKey & dpf, IntegralT from_node,
|
|||
IntegralT to_node, OutputBuffers && outbufs, IntervalMemoizer && memoizer,
|
||||
std::index_sequence<IIs...> idxs, std::size_t start = 0)
|
||||
{
|
||||
using node_type = typename DpfKey::exterior_node;
|
||||
using outputs_tuple = typename DpfKey::concrete_outputs_tuple;
|
||||
HEDLEY_PRAGMA(GCC diagnostic push)
|
||||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
||||
using range = leaf_prg_range<node_type, outputs_tuple, Is...>;
|
||||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||||
if constexpr (range::is_contiguous)
|
||||
// Fused exterior needs classic leaf packing (`concrete_outputs_tuple` +
|
||||
// contiguous PRG lanes). Multi-level / cmp keys use the per-slot walk.
|
||||
// Extractable keys stretch leaves with `extractable_leaf_prg` (leaf XOF),
|
||||
// not `exterior_prg`. The fused path expands with AES and breaks the
|
||||
// programmed packed leaf (cold opens still cancel; hot lanes do not).
|
||||
if constexpr (!is_multilevel_key_v<DpfKey> && !DpfKey::is_extractable)
|
||||
{
|
||||
eval_interval_exterior_fused<Is...>(dpf, from_node, to_node, outbufs,
|
||||
memoizer, idxs, start);
|
||||
}
|
||||
else
|
||||
{
|
||||
(eval_interval_exterior<Is>(dpf, from_node, to_node,
|
||||
utils::get<IIs>(outbufs), memoizer, start), ...);
|
||||
using node_type = typename DpfKey::exterior_node;
|
||||
using outputs_tuple = typename DpfKey::concrete_outputs_tuple;
|
||||
HEDLEY_PRAGMA(GCC diagnostic push)
|
||||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
||||
using range = leaf_prg_range<node_type, outputs_tuple, Is...>;
|
||||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||||
if constexpr (range::is_contiguous)
|
||||
{
|
||||
eval_interval_exterior_fused<Is...>(dpf, from_node, to_node, outbufs,
|
||||
memoizer, idxs, start);
|
||||
return;
|
||||
}
|
||||
}
|
||||
(eval_interval_exterior<Is>(dpf, from_node, to_node,
|
||||
utils::get<IIs>(outbufs), memoizer, start), ...);
|
||||
}
|
||||
|
||||
template <std::size_t ...Is,
|
||||
|
|
@ -356,7 +426,7 @@ template <std::size_t ...Is,
|
|||
std::size_t ...IIs>
|
||||
auto eval_interval_impl(const DpfKey & dpf, InputT from, InputT to,
|
||||
OutputBuffers && outbufs, IntervalMemoizer && memoizer,
|
||||
std::index_sequence<IIs...>)
|
||||
std::index_sequence<IIs...>, proof_token * pi = nullptr)
|
||||
{
|
||||
using dpf_type = DpfKey;
|
||||
using integral_type = typename DpfKey::integral_type;
|
||||
|
|
@ -378,7 +448,8 @@ auto eval_interval_impl(const DpfKey & dpf, InputT from, InputT to,
|
|||
for (std::size_t s = 0; s < segs.n; ++s)
|
||||
{
|
||||
const auto & seg = segs.seg[s];
|
||||
internal::eval_interval_interior(dpf, seg.from_node, seg.to_node, memoizer);
|
||||
internal::eval_interval_interior(dpf, seg.from_node, seg.to_node, memoizer,
|
||||
DpfKey::depth, pi);
|
||||
eval_interval_exterior_all<Is...>(dpf, seg.from_node, seg.to_node, outbufs,
|
||||
memoizer, idxs, start);
|
||||
start += seg.count;
|
||||
|
|
@ -393,14 +464,15 @@ template <std::size_t ...Is,
|
|||
std::size_t ...IIs>
|
||||
auto eval_interval(const DpfKey & dpf, InputT from, InputT to,
|
||||
OutputBuffers && outbufs, IntervalMemoizer && memoizer,
|
||||
std::index_sequence<IIs...>)
|
||||
std::index_sequence<IIs...>, proof_token * pi = nullptr)
|
||||
{
|
||||
using dpf_type = DpfKey;
|
||||
constexpr auto mod_pow_2 = utils::mod_pow_2<InputT>{};
|
||||
constexpr auto to_integral_t = utils::to_integral_type<InputT>{};
|
||||
constexpr auto bits = utils::bitlength_of_v<InputT>;
|
||||
|
||||
eval_interval_impl<Is...>(dpf, from, to, outbufs, memoizer, std::make_index_sequence<sizeof...(Is)>());
|
||||
eval_interval_impl<Is...>(dpf, from, to, outbufs, memoizer,
|
||||
std::make_index_sequence<sizeof...(Is)>(), pi);
|
||||
|
||||
// `to_integral_type` widens to at least `size_t`. Subtracting in that
|
||||
// wider type loses wrap-around of a narrower input domain (e.g. int16
|
||||
|
|
@ -439,7 +511,9 @@ auto eval_interval(const DpfKey & dpf, InputT from, InputT to,
|
|||
/// @param outbufs named buffer, or a tuple of buffers when several outputs
|
||||
/// are selected. Must outlive the returned iterable.
|
||||
/// @param memoizer workspace sized for at least this interval
|
||||
/// @param pi proof token folded along the interval, or null
|
||||
/// @return an iterable over the written outputs
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
|
|
@ -449,11 +523,60 @@ template <std::size_t I = 0,
|
|||
std::enable_if_t<looks_like_dpf_key_v<DpfKey> && !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto eval_interval(const DpfKey & dpf, InputT from, InputT to,
|
||||
OutputBuffers & outbufs, IntervalMemoizer && memoizer) // NOLINT(runtime/references)
|
||||
OutputBuffers & outbufs, IntervalMemoizer && memoizer, // NOLINT(runtime/references)
|
||||
proof_token * pi = nullptr)
|
||||
{
|
||||
assert_not_wildcard_output<I, Is...>(dpf);
|
||||
|
||||
return internal::eval_interval<I, Is...>(dpf, dpf.offset_x(from), dpf.offset_x(to), outbufs, memoizer, std::make_index_sequence<1+sizeof...(Is)>());
|
||||
return internal::eval_interval<I, Is...>(dpf, dpf.offset_x(from), dpf.offset_x(to),
|
||||
outbufs, memoizer, std::make_index_sequence<1+sizeof...(Is)>(), pi);
|
||||
}
|
||||
|
||||
/// @brief Evaluate `[from, to]` and fold a once-per-BFS-node VDPF proof.
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
typename InputT,
|
||||
typename OutputBuffers,
|
||||
typename IntervalMemoizer = dpf::basic_interval_memoizer<DpfKey>,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey> && !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto eval_interval(const DpfKey & dpf, InputT from, InputT to,
|
||||
OutputBuffers & outbufs, IntervalMemoizer && memoizer, prove_ref pr) // NOLINT(runtime/references)
|
||||
{
|
||||
static_assert(DpfKey::is_verifiable,
|
||||
"eval_interval(..., prove(π)): key must carry dpf::verifiable");
|
||||
detail::vdpf::init_proof(pr.token, dpf);
|
||||
auto out = eval_interval<I, Is...>(dpf, from, to, outbufs,
|
||||
std::forward<IntervalMemoizer>(memoizer), &pr.token);
|
||||
detail::vdpf::fold_output_binding(pr.token, dpf);
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief Evaluate `[from, to]` and fold each written output into a sketch.
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
typename InputT,
|
||||
typename OutputBuffers,
|
||||
typename IntervalMemoizer = dpf::basic_interval_memoizer<DpfKey>,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey> && !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto eval_interval(const DpfKey & dpf, InputT from, InputT to,
|
||||
OutputBuffers & outbufs, IntervalMemoizer && memoizer, sketch_ref & sk) // NOLINT(runtime/references)
|
||||
{
|
||||
static_assert(DpfKey::is_extractable,
|
||||
"eval_interval(..., sketch(σ)): key must carry dpf::extractable");
|
||||
auto ret = eval_interval<I, Is...>(dpf, from, to, outbufs,
|
||||
std::forward<IntervalMemoizer>(memoizer));
|
||||
if constexpr (sizeof...(Is) == 0)
|
||||
{
|
||||
for (std::size_t k = 0; k < utils::size(outbufs); ++k)
|
||||
sk.absorb(outbufs[k]);
|
||||
}
|
||||
return ret;
|
||||
}
|
||||
|
||||
/// @brief Evaluate `[from, to]` into `outbufs`, allocating a basic interval memoizer.
|
||||
|
|
@ -463,6 +586,7 @@ auto eval_interval(const DpfKey & dpf, InputT from, InputT to,
|
|||
/// @param to the inclusive end of the range
|
||||
/// @param outbufs the named output buffers
|
||||
/// @return an iterable over the written outputs
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
|
|
@ -477,7 +601,28 @@ auto eval_interval(const DpfKey & dpf, InputT from, InputT to,
|
|||
OutputBuffers & outbufs) // NOLINT(runtime/references)
|
||||
{
|
||||
return eval_interval<I, Is...>(dpf, from, to, outbufs,
|
||||
dpf::make_basic_interval_memoizer<DpfKey>(from, to));
|
||||
dpf::make_basic_interval_memoizer(dpf, from, to));
|
||||
}
|
||||
|
||||
/// @brief Evaluate `[from, to]` into `outbufs` and fold a VDPF proof.
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
typename InputT,
|
||||
typename OutputBuffers,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey> && !is_multilevel_key_v<DpfKey>, bool> = true,
|
||||
std::enable_if_t<!std::is_base_of_v<
|
||||
dpf::interval_memoizer_base<unwrap_party_key_t<DpfKey>>,
|
||||
std::decay_t<OutputBuffers>>, bool> = true>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto eval_interval(const DpfKey & dpf, InputT from, InputT to,
|
||||
OutputBuffers & outbufs, prove_ref pr) // NOLINT(runtime/references)
|
||||
{
|
||||
static_assert(DpfKey::is_verifiable,
|
||||
"eval_interval(..., prove(π)): key must carry dpf::verifiable");
|
||||
return eval_interval<I, Is...>(dpf, from, to, outbufs,
|
||||
dpf::make_basic_interval_memoizer(dpf, from, to), pr);
|
||||
}
|
||||
|
||||
/// @brief Evaluate `[from, to]` with a caller-supplied memoizer.
|
||||
|
|
@ -488,6 +633,7 @@ auto eval_interval(const DpfKey & dpf, InputT from, InputT to,
|
|||
/// @param memoizer the memoizer built for this key
|
||||
/// @return `std::pair` of a new buffer (or tuple of buffers) and an iterable
|
||||
/// into that buffer.
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
|
|
@ -512,9 +658,34 @@ auto eval_interval(const DpfKey & dpf, InputT from, InputT to,
|
|||
return std::make_pair(std::move(outbufs), std::move(iterable));
|
||||
}
|
||||
|
||||
/// @brief Evaluate `[from, to]` with a memoizer and fold a VDPF proof.
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
typename InputT,
|
||||
typename IntervalMemoizer,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey> && !is_multilevel_key_v<DpfKey>, bool> = true,
|
||||
std::enable_if_t<std::is_base_of_v<
|
||||
dpf::interval_memoizer_base<unwrap_party_key_t<DpfKey>>,
|
||||
std::decay_t<IntervalMemoizer>>, bool> = true>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto eval_interval(const DpfKey & dpf, InputT from, InputT to,
|
||||
IntervalMemoizer && memoizer, prove_ref pr)
|
||||
{
|
||||
static_assert(DpfKey::is_verifiable,
|
||||
"eval_interval(..., prove(π)): key must carry dpf::verifiable");
|
||||
auto outbufs = utils::make_tuple(
|
||||
make_output_buffer_for_interval<I>(dpf, from, to),
|
||||
make_output_buffer_for_interval<Is>(dpf, from, to)...);
|
||||
auto iterable = eval_interval<I, Is...>(dpf, from, to, outbufs, memoizer, pr);
|
||||
return std::make_pair(std::move(outbufs), std::move(iterable));
|
||||
}
|
||||
|
||||
/// @brief Evaluate `[from, to]`, allocating a basic interval memoizer and a buffer.
|
||||
/// @return `std::pair` of a new buffer (or tuple of buffers) and an iterable
|
||||
/// into that buffer.
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
|
|
@ -524,7 +695,218 @@ HEDLEY_ALWAYS_INLINE
|
|||
auto eval_interval(const DpfKey & dpf, InputT from, InputT to)
|
||||
{
|
||||
return eval_interval<I, Is...>(dpf, from, to,
|
||||
dpf::make_basic_interval_memoizer<DpfKey>(from, to));
|
||||
dpf::make_basic_interval_memoizer(dpf, from, to));
|
||||
}
|
||||
|
||||
/// @brief Fold every truncated-tree node of `[from, to]` into `pi`.
|
||||
/// @details Once per BFS node — same absorption as `eval_interval(..., prove(π))`.
|
||||
/// Caller must `init_proof` first, or use `prove_interval` below.
|
||||
template <typename KeyT, typename InputT>
|
||||
void prove_fold_interval(const KeyT & key, InputT from, InputT to,
|
||||
proof_token & pi)
|
||||
{
|
||||
static_assert(KeyT::is_verifiable,
|
||||
"prove_fold_interval: key must carry dpf::verifiable");
|
||||
using dpf_type = KeyT;
|
||||
using input_type = typename KeyT::input_type;
|
||||
using integral_type = typename KeyT::integral_type;
|
||||
|
||||
auto from_x = key.offset_x(static_cast<input_type>(from));
|
||||
auto to_x = key.offset_x(static_cast<input_type>(to));
|
||||
utils::flip_msb_if_signed_integral(from_x);
|
||||
utils::flip_msb_if_signed_integral(to_x);
|
||||
|
||||
integral_type from_node = utils::get_from_node<dpf_type>(from_x);
|
||||
integral_type to_node = utils::get_to_node<dpf_type>(to_x);
|
||||
constexpr auto to_int = utils::to_integral_type<input_type>{};
|
||||
const bool wraps = utils::interval_wraps(
|
||||
static_cast<integral_type>(to_int(from_x)),
|
||||
static_cast<integral_type>(to_int(to_x)),
|
||||
utils::bitlength_of_v<input_type>);
|
||||
auto segs = utils::split_leaf_nodes(from_node, to_node, key.depth, wraps);
|
||||
auto memo = make_basic_interval_memoizer(key, from, to);
|
||||
for (std::size_t s = 0; s < segs.n; ++s)
|
||||
{
|
||||
const auto & seg = segs.seg[s];
|
||||
internal::eval_interval_interior(key, seg.from_node, seg.to_node, memo,
|
||||
KeyT::depth, &pi);
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Initialise `pr.token` and fold `[from, to]` once per BFS node.
|
||||
/// @tparam KeyT verifiable key
|
||||
/// @tparam InputT input domain type
|
||||
/// @param key the party key
|
||||
/// @param from inclusive start
|
||||
/// @param to inclusive end
|
||||
/// @param pr proof token replaced with the interval fold
|
||||
template <typename KeyT, typename InputT>
|
||||
void prove_interval(const KeyT & key, InputT from, InputT to, prove_ref pr)
|
||||
{
|
||||
detail::vdpf::init_proof(pr.token, key);
|
||||
prove_fold_interval(key, from, to, pr.token);
|
||||
detail::vdpf::fold_output_binding(pr.token, key);
|
||||
}
|
||||
|
||||
/// @brief Initialise `pr.token` and fold the full domain once per BFS node.
|
||||
/// @tparam KeyT verifiable key
|
||||
/// @param key the party key
|
||||
/// @param pr proof token replaced with the full-domain fold
|
||||
template <typename KeyT>
|
||||
void prove_full(const KeyT & key, prove_ref pr)
|
||||
{
|
||||
static_assert(KeyT::is_verifiable,
|
||||
"prove_full: key must carry dpf::verifiable");
|
||||
using input_type = typename KeyT::input_type;
|
||||
prove_interval(key, std::numeric_limits<input_type>::min(),
|
||||
std::numeric_limits<input_type>::max(), pr);
|
||||
}
|
||||
|
||||
/// @}
|
||||
|
||||
namespace internal
|
||||
{
|
||||
|
||||
template <std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
typename InputT,
|
||||
typename OutputBuffers,
|
||||
typename IntervalMemoizer,
|
||||
std::size_t ...IIs>
|
||||
auto defer_eval_interval(const DpfKey & dpf, InputT from, InputT to,
|
||||
OutputBuffers & outbufs, IntervalMemoizer && memoizer,
|
||||
std::index_sequence<IIs...>)
|
||||
{
|
||||
using dpf_type = DpfKey;
|
||||
using input_type = typename dpf_type::input_type;
|
||||
|
||||
const auto min = std::numeric_limits<input_type>::min();
|
||||
const auto max = std::numeric_limits<input_type>::max();
|
||||
|
||||
// Full-domain identity traversal (offset unknown). Same interior/exterior
|
||||
// path as eager eval over `[min, max]` without folding `offset_x`.
|
||||
eval_interval_impl<Is...>(dpf, min, max, outbufs, memoizer,
|
||||
std::make_index_sequence<sizeof...(Is)>());
|
||||
|
||||
return utils::make_tuple(
|
||||
deferred_rotated_subinterval(dpf,
|
||||
std::begin(utils::get<IIs>(outbufs)),
|
||||
std::end(utils::get<IIs>(outbufs)),
|
||||
from, to,
|
||||
dpf_type::outputs_per_leaf)...);
|
||||
}
|
||||
|
||||
} // namespace internal
|
||||
|
||||
/// @name Deferred (pre-assign) evaluation
|
||||
/// @{
|
||||
|
||||
/// @brief Full-domain eval while the input offset is still unset.
|
||||
/// @details Requires a wildcard input that is not yet ready, and assigned
|
||||
/// leaf outputs `I, Is...`. `outbufs` must be sized for the **full**
|
||||
/// input domain (`make_output_buffer_for_full`). After
|
||||
/// `assign_wildcard_input`, call `.get()` on each returned view.
|
||||
/// @return one `deferred_rotated_subinterval` per selected output
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
typename InputT,
|
||||
typename OutputBuffers,
|
||||
typename IntervalMemoizer,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
auto defer_eval_interval(const DpfKey & dpf, InputT from, InputT to,
|
||||
OutputBuffers & outbufs, IntervalMemoizer && memoizer) // NOLINT(runtime/references)
|
||||
{
|
||||
static_assert(is_wildcard_v<typename DpfKey::raw_input_type>,
|
||||
"defer_eval_interval: key input must be a wildcard_value");
|
||||
assert_wildcard_input(dpf);
|
||||
assert_not_wildcard_output<I, Is...>(dpf);
|
||||
|
||||
return internal::defer_eval_interval<I, Is...>(dpf, from, to, outbufs,
|
||||
std::forward<IntervalMemoizer>(memoizer),
|
||||
std::make_index_sequence<1 + sizeof...(Is)>());
|
||||
}
|
||||
|
||||
/// @brief `defer_eval_interval` with a basic full-domain memoizer.
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
typename InputT,
|
||||
typename OutputBuffers,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>, bool> = true,
|
||||
std::enable_if_t<!std::is_base_of_v<
|
||||
dpf::interval_memoizer_base<unwrap_party_key_t<DpfKey>>,
|
||||
std::decay_t<OutputBuffers>>, bool> = true>
|
||||
auto defer_eval_interval(const DpfKey & dpf, InputT from, InputT to,
|
||||
OutputBuffers & outbufs) // NOLINT(runtime/references)
|
||||
{
|
||||
return defer_eval_interval<I, Is...>(dpf, from, to, outbufs,
|
||||
dpf::make_basic_full_memoizer(dpf));
|
||||
}
|
||||
|
||||
/// @brief Interior-only traverse of `[offset_x(from), offset_x(to)]`.
|
||||
/// @details Requires an assigned input. Skips exterior so the leaf may still
|
||||
/// be a wildcard; finish with `eval_interval` / exterior once the
|
||||
/// leaf is assigned (memoizer retains the interior).
|
||||
template <typename DpfKey,
|
||||
typename InputT,
|
||||
typename IntervalMemoizer,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
void defer_traverse_interval(const DpfKey & dpf, InputT from, InputT to,
|
||||
IntervalMemoizer & memoizer) // NOLINT(runtime/references)
|
||||
{
|
||||
assert_not_wildcard_input(dpf);
|
||||
|
||||
using dpf_type = DpfKey;
|
||||
using input_type = typename dpf_type::input_type;
|
||||
using integral_type = typename dpf_type::integral_type;
|
||||
|
||||
auto tfrom = dpf.offset_x(from);
|
||||
auto tto = dpf.offset_x(to);
|
||||
utils::flip_msb_if_signed_integral(tfrom);
|
||||
utils::flip_msb_if_signed_integral(tto);
|
||||
|
||||
integral_type from_node = utils::get_from_node<dpf_type>(tfrom);
|
||||
integral_type to_node = utils::get_to_node<dpf_type>(tto);
|
||||
constexpr auto to_int = utils::to_integral_type<input_type>{};
|
||||
const bool wraps = utils::interval_wraps(
|
||||
static_cast<integral_type>(to_int(tfrom)),
|
||||
static_cast<integral_type>(to_int(tto)),
|
||||
utils::bitlength_of_v<input_type>);
|
||||
auto segs = utils::split_leaf_nodes(from_node, to_node, dpf.depth, wraps);
|
||||
for (std::size_t s = 0; s < segs.n; ++s)
|
||||
{
|
||||
const auto & seg = segs.seg[s];
|
||||
internal::eval_interval_interior(dpf, seg.from_node, seg.to_node,
|
||||
memoizer);
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Interior-only full-domain traverse (input and/or leaf may be unset).
|
||||
/// @details Fills the memoizer for every interior node. Complete exterior
|
||||
/// (and any input rotation) after the missing wildcards are assigned.
|
||||
template <typename DpfKey,
|
||||
typename IntervalMemoizer,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
void defer_traverse_full(const DpfKey & dpf,
|
||||
IntervalMemoizer & memoizer) // NOLINT(runtime/references)
|
||||
{
|
||||
using dpf_type = DpfKey;
|
||||
using input_type = typename dpf_type::input_type;
|
||||
using integral_type = typename dpf_type::integral_type;
|
||||
|
||||
auto from = std::numeric_limits<input_type>::min();
|
||||
auto to = std::numeric_limits<input_type>::max();
|
||||
utils::flip_msb_if_signed_integral(from);
|
||||
utils::flip_msb_if_signed_integral(to);
|
||||
|
||||
integral_type from_node = utils::get_from_node<dpf_type>(from);
|
||||
integral_type to_node = utils::get_to_node<dpf_type>(to);
|
||||
internal::eval_interval_interior(dpf, from_node, to_node, memoizer);
|
||||
}
|
||||
|
||||
/// @}
|
||||
|
|
|
|||
155
include/dpf/eval_peel.hpp
Normal file
155
include/dpf/eval_peel.hpp
Normal file
|
|
@ -0,0 +1,155 @@
|
|||
/// @file dpf/eval_peel.hpp
|
||||
/// @brief Eval overloads that accept a wrapper and evaluate its `dpf_key`.
|
||||
/// @details Include this after the eval headers. A call passes through while
|
||||
/// the argument has a public `dpf_key` member and is not itself a
|
||||
/// DPF key. `dpf3` / `dpf3_cmp` / `dpf3_ic` keep their key-first
|
||||
/// protocol overloads; `eval_*(cmp, …)` and `eval_*(out<I>, …)` still
|
||||
/// peel those wrappers down to the inner DPF key.
|
||||
/// @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_DPF_EVAL_PEEL_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_EVAL_PEEL_HPP__
|
||||
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
|
||||
#include "dpf/eval_target.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
template <typename T>
|
||||
inline constexpr bool peel_key_v =
|
||||
has_embedded_dpf_key_v<std::decay_t<T>>
|
||||
&& !owns_protocol_eval_v<std::decay_t<T>>;
|
||||
|
||||
template <typename T>
|
||||
inline constexpr bool peel_tag_key_v =
|
||||
has_embedded_dpf_key_v<std::decay_t<T>>;
|
||||
|
||||
#define LIBDPF_PEEL_KEY_FIRST(fn) \
|
||||
template <std::size_t I = 0, \
|
||||
typename KeyT, \
|
||||
typename... Args, \
|
||||
std::enable_if_t<peel_key_v<KeyT>, int> = 0> \
|
||||
decltype(auto) fn(const KeyT & key, Args &&... args) \
|
||||
{ \
|
||||
return fn<I>(key.dpf_key, std::forward<Args>(args)...); \
|
||||
}
|
||||
|
||||
#define LIBDPF_PEEL_TAG_FIRST(fn) \
|
||||
template <typename Tag, \
|
||||
typename KeyT, \
|
||||
typename... Args, \
|
||||
std::enable_if_t<peel_tag_key_v<KeyT> \
|
||||
&& is_eval_channel_tag_v<std::decay_t<Tag>>, int> = 0> \
|
||||
decltype(auto) fn(Tag tag, const KeyT & key, Args &&... args) \
|
||||
{ \
|
||||
return fn(tag, key.dpf_key, std::forward<Args>(args)...); \
|
||||
}
|
||||
|
||||
LIBDPF_PEEL_KEY_FIRST(eval_point)
|
||||
LIBDPF_PEEL_KEY_FIRST(eval_interval)
|
||||
LIBDPF_PEEL_KEY_FIRST(eval_full)
|
||||
LIBDPF_PEEL_KEY_FIRST(eval_sequence)
|
||||
LIBDPF_PEEL_KEY_FIRST(eval_inner_product)
|
||||
LIBDPF_PEEL_KEY_FIRST(eval_full_inner_product)
|
||||
LIBDPF_PEEL_KEY_FIRST(eval_sequence_inner_product)
|
||||
|
||||
LIBDPF_PEEL_TAG_FIRST(eval_point)
|
||||
LIBDPF_PEEL_TAG_FIRST(eval_interval)
|
||||
LIBDPF_PEEL_TAG_FIRST(eval_full)
|
||||
LIBDPF_PEEL_TAG_FIRST(eval_sequence)
|
||||
LIBDPF_PEEL_TAG_FIRST(eval_sequence_breadth_first)
|
||||
LIBDPF_PEEL_TAG_FIRST(eval_inner_product)
|
||||
LIBDPF_PEEL_TAG_FIRST(eval_full_inner_product)
|
||||
LIBDPF_PEEL_TAG_FIRST(make_output_buffer)
|
||||
LIBDPF_PEEL_TAG_FIRST(make_sequence_recipe)
|
||||
LIBDPF_PEEL_TAG_FIRST(eval_prefixes)
|
||||
LIBDPF_PEEL_TAG_FIRST(eval_prefix_inner_product)
|
||||
|
||||
#undef LIBDPF_PEEL_KEY_FIRST
|
||||
#undef LIBDPF_PEEL_TAG_FIRST
|
||||
|
||||
template <typename KeyT,
|
||||
typename... Args,
|
||||
std::enable_if_t<peel_key_v<KeyT>, int> = 0>
|
||||
decltype(auto) eval_sequence_xor(const KeyT & key, Args &&... args)
|
||||
{
|
||||
return eval_sequence_xor(key.dpf_key, std::forward<Args>(args)...);
|
||||
}
|
||||
|
||||
template <typename KeyT,
|
||||
typename... Args,
|
||||
std::enable_if_t<peel_key_v<KeyT>, int> = 0>
|
||||
decltype(auto) prove_cmp_interval(const KeyT & key, Args &&... args)
|
||||
{
|
||||
return prove_cmp_interval(key.dpf_key, std::forward<Args>(args)...);
|
||||
}
|
||||
|
||||
template <typename KeyT,
|
||||
typename... Args,
|
||||
std::enable_if_t<peel_key_v<KeyT>, int> = 0>
|
||||
decltype(auto) prove_cmp_full(const KeyT & key, Args &&... args)
|
||||
{
|
||||
return prove_cmp_full(key.dpf_key, std::forward<Args>(args)...);
|
||||
}
|
||||
|
||||
template <typename KeyT,
|
||||
typename... Args,
|
||||
std::enable_if_t<peel_key_v<KeyT>, int> = 0>
|
||||
decltype(auto) prove_cmp_sequence(const KeyT & key, Args &&... args)
|
||||
{
|
||||
return prove_cmp_sequence(key.dpf_key, std::forward<Args>(args)...);
|
||||
}
|
||||
|
||||
template <std::size_t I = 0,
|
||||
typename Buffer,
|
||||
typename KeyT,
|
||||
typename... Args,
|
||||
std::enable_if_t<peel_key_v<KeyT>, int> = 0>
|
||||
void eval_full_add_into(Buffer & buf, const KeyT & key, Args &&... args)
|
||||
{
|
||||
eval_full_add_into<I>(buf, key.dpf_key, std::forward<Args>(args)...);
|
||||
}
|
||||
|
||||
/// @brief Row-wise and column inner products keep the tag in front of the key.
|
||||
template <typename Tag,
|
||||
typename KeyT,
|
||||
typename... Args,
|
||||
std::enable_if_t<peel_tag_key_v<KeyT>
|
||||
&& (std::is_same_v<std::decay_t<Tag>, paired_t>
|
||||
|| std::is_same_v<std::decay_t<Tag>, columns_t>), int> = 0>
|
||||
decltype(auto) eval_inner_product(Tag tag, const KeyT & key, Args &&... args)
|
||||
{
|
||||
return eval_inner_product(tag, key.dpf_key, std::forward<Args>(args)...);
|
||||
}
|
||||
|
||||
template <typename Tag,
|
||||
typename KeyT,
|
||||
typename... Args,
|
||||
std::enable_if_t<peel_tag_key_v<KeyT>
|
||||
&& (std::is_same_v<std::decay_t<Tag>, paired_t>
|
||||
|| std::is_same_v<std::decay_t<Tag>, columns_t>), int> = 0>
|
||||
decltype(auto) eval_full_inner_product(Tag tag, const KeyT & key, Args &&... args)
|
||||
{
|
||||
return eval_full_inner_product(tag, key.dpf_key, std::forward<Args>(args)...);
|
||||
}
|
||||
|
||||
template <typename Tag,
|
||||
typename KeyT,
|
||||
typename... Args,
|
||||
std::enable_if_t<peel_tag_key_v<KeyT>
|
||||
&& (std::is_same_v<std::decay_t<Tag>, paired_t>
|
||||
|| std::is_same_v<std::decay_t<Tag>, columns_t>), int> = 0>
|
||||
decltype(auto) eval_sequence_inner_product(Tag tag, const KeyT & key,
|
||||
Args &&... args)
|
||||
{
|
||||
return eval_sequence_inner_product(tag, key.dpf_key,
|
||||
std::forward<Args>(args)...);
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_EVAL_PEEL_HPP__
|
||||
|
|
@ -44,25 +44,31 @@ inline auto eval_point_interior(const DpfKey & dpf, InputT && x, PathMemoizer &&
|
|||
|
||||
auto level_index = detail::path_resume_for_level(path, dpf, x, dpf.depth);
|
||||
|
||||
DPF_UNROLL_LOOP
|
||||
for (auto mask = dpf.msb_mask>>(level_index-1);
|
||||
level_index <= dpf.depth; ++level_index, mask>>=1)
|
||||
// Same guard as `ensure_level`: skip the `msb_mask` shift when already done.
|
||||
if (level_index <= dpf.depth)
|
||||
{
|
||||
bool bit = !!(mask & x);
|
||||
auto cw = dpf.correction_word(level_index-1, bit);
|
||||
const bool is_last = dpf_type::tree::is_last_level(level_index - 1,
|
||||
dpf.depth);
|
||||
path[level_index] = dpf_type::traverse_interior(path[level_index-1],
|
||||
cw, bit, is_last);
|
||||
if constexpr (dpf_type::is_verifiable)
|
||||
DPF_UNROLL_LOOP
|
||||
for (auto mask = dpf.msb_mask>>(level_index-1);
|
||||
level_index <= dpf.depth; ++level_index, mask>>=1)
|
||||
{
|
||||
if (pi != nullptr)
|
||||
bool bit = !!(mask & x);
|
||||
auto cw = dpf.correction_word(level_index-1, bit);
|
||||
const bool is_last = dpf_type::tree::is_last_level(level_index - 1,
|
||||
dpf.depth);
|
||||
path[level_index] = dpf_type::traverse_interior(path[level_index-1],
|
||||
cw, bit, is_last);
|
||||
if constexpr (dpf_type::is_verifiable)
|
||||
{
|
||||
const auto x_bits = static_cast<psnip_uint64_t>(
|
||||
utils::to_integral_type<std::decay_t<InputT>>{}(x)
|
||||
>> (utils::bitlength_of_v<std::decay_t<InputT>> - level_index));
|
||||
detail::vdpf::fold_node(*pi, level_index - 1, x_bits,
|
||||
path[level_index], dpf.correction_seeds()[level_index - 1]);
|
||||
if (pi != nullptr)
|
||||
{
|
||||
const auto x_bits = static_cast<psnip_uint64_t>(
|
||||
utils::to_integral_type<std::decay_t<InputT>>{}(x)
|
||||
>> (utils::bitlength_of_v<std::decay_t<InputT>>
|
||||
- level_index));
|
||||
detail::vdpf::fold_node(*pi, level_index - 1, x_bits,
|
||||
path[level_index],
|
||||
dpf.correction_seeds()[level_index - 1]);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -96,6 +102,7 @@ auto eval_point(const DpfKey & dpf, InputT && x, PathMemoizer && path,
|
|||
} // namespace internal
|
||||
|
||||
/// Evaluate output `I` at `x`.
|
||||
/// \complexity O(n) time. n is `depth`. One interior traversal per level from the memoizer resume index through the leaf. Extra space is the path memoizer (O(n) nodes, or one node if it does not memoize). A proof token adds one fold per level walked.
|
||||
template <std::size_t I = 0,
|
||||
typename DpfKey,
|
||||
typename InputT,
|
||||
|
|
@ -112,7 +119,33 @@ auto eval_point(const DpfKey & dpf, InputT && x, PathMemoizer && path = PathMemo
|
|||
internal::eval_point<I>(dpf, tx, path), tx);
|
||||
}
|
||||
|
||||
/// Evaluate at `x`, folding newly walked nodes into an existing proof token.
|
||||
/// @details Does not `init_proof`; used by bulk sequence evals that share a path.
|
||||
/// \complexity O(n) time. n is `depth`. One interior traversal per level from the memoizer resume index through the leaf. Extra space is the path memoizer (O(n) nodes, or one node if it does not memoize). A proof token adds one fold per level walked.
|
||||
template <std::size_t I = 0,
|
||||
typename DpfKey,
|
||||
typename InputT,
|
||||
typename PathMemoizer,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey> && !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto eval_point(const DpfKey & dpf, InputT && x, PathMemoizer && path,
|
||||
proof_token * pi)
|
||||
{
|
||||
assert_not_wildcard_output<I>(dpf);
|
||||
using output_type = typename DpfKey::concrete_output_type<I>;
|
||||
|
||||
auto tx = dpf.offset_x(x);
|
||||
return make_eval_dpf_output<DpfKey, output_type>(
|
||||
internal::eval_point<I>(dpf, tx, path, pi), tx);
|
||||
}
|
||||
|
||||
/// Evaluate and fold a VDPF proof token for the walked path.
|
||||
/// @details A fresh token always folds every node on the path, including
|
||||
/// nodes already stored in a warm path memoizer (hash the cached
|
||||
/// seeds; do not re-expand the PRG). Callers that continue one
|
||||
/// accumulator across many points use the `proof_token *` overload
|
||||
/// without `init_proof`.
|
||||
/// \complexity O(n) time. n is `depth`. One interior traversal per level from the memoizer resume index through the leaf. Extra space is the path memoizer (O(n) nodes, or one node if it does not memoize). A proof token adds one fold per level walked.
|
||||
template <std::size_t I = 0,
|
||||
typename DpfKey,
|
||||
typename InputT,
|
||||
|
|
@ -124,16 +157,68 @@ auto eval_point(const DpfKey & dpf, InputT && x, prove_ref pr,
|
|||
{
|
||||
static_assert(DpfKey::is_verifiable,
|
||||
"eval_point(..., prove(π)): key must carry dpf::verifiable");
|
||||
assert_not_wildcard_output<I>(dpf);
|
||||
using output_type = typename DpfKey::concrete_output_type<I>;
|
||||
|
||||
detail::vdpf::init_proof(pr.token, dpf);
|
||||
auto tx = dpf.offset_x(x);
|
||||
return make_eval_dpf_output<DpfKey, output_type>(
|
||||
internal::eval_point<I>(dpf, tx, path, &pr.token), tx);
|
||||
auto walk_x = dpf.offset_x(x);
|
||||
utils::flip_msb_if_signed_integral(walk_x);
|
||||
if constexpr (DpfKey::is_verifiable)
|
||||
{
|
||||
const auto resume = detail::path_resume_for_level(
|
||||
path, dpf, walk_x, dpf.depth);
|
||||
for (std::size_t level_index = 1; level_index < resume; ++level_index)
|
||||
{
|
||||
const auto x_bits = static_cast<psnip_uint64_t>(
|
||||
utils::to_integral_type<std::decay_t<decltype(walk_x)>>{}(
|
||||
walk_x)
|
||||
>> (utils::bitlength_of_v<std::decay_t<decltype(walk_x)>>
|
||||
- level_index));
|
||||
detail::vdpf::fold_node(pr.token, level_index - 1, x_bits,
|
||||
path[level_index],
|
||||
dpf.correction_seeds()[level_index - 1]);
|
||||
}
|
||||
}
|
||||
auto out = eval_point<I>(dpf, std::forward<InputT>(x),
|
||||
std::forward<PathMemoizer>(path), &pr.token);
|
||||
detail::vdpf::fold_output_binding(pr.token, dpf);
|
||||
return out;
|
||||
}
|
||||
|
||||
/// Evaluate and fold the output into a weight-1 sketch.
|
||||
/// \complexity O(n) time. n is `depth`. One interior traversal per level from the memoizer resume index through the leaf. Extra space is the path memoizer (O(n) nodes, or one node if it does not memoize). A proof token adds one fold per level walked.
|
||||
template <std::size_t I = 0,
|
||||
typename DpfKey,
|
||||
typename InputT,
|
||||
typename PathMemoizer = dpf::nonmemoizing_path_memoizer<DpfKey>,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>, bool> = true>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto eval_point(const DpfKey & dpf, InputT && x, sketch_ref & sk,
|
||||
PathMemoizer && path = PathMemoizer{})
|
||||
{
|
||||
static_assert(DpfKey::is_extractable,
|
||||
"eval_point(..., sketch(σ)): key must carry dpf::extractable");
|
||||
auto out = eval_point<I>(dpf, std::forward<InputT>(x),
|
||||
std::forward<PathMemoizer>(path));
|
||||
if constexpr (DpfKey::is_extractable)
|
||||
sk.absorb((*out).raw());
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief Rvalue convenience for a one-shot `sketch(local, rs)` temporary.
|
||||
/// \complexity O(n) time. n is `depth`. One interior traversal per level from the memoizer resume index through the leaf. Extra space is the path memoizer (O(n) nodes, or one node if it does not memoize). A proof token adds one fold per level walked.
|
||||
template <std::size_t I = 0,
|
||||
typename DpfKey,
|
||||
typename InputT,
|
||||
typename PathMemoizer = dpf::nonmemoizing_path_memoizer<DpfKey>,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>, bool> = true>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto eval_point(const DpfKey & dpf, InputT && x, sketch_ref && sk,
|
||||
PathMemoizer && path = PathMemoizer{})
|
||||
{
|
||||
return eval_point<I>(dpf, std::forward<InputT>(x), sk,
|
||||
std::forward<PathMemoizer>(path));
|
||||
}
|
||||
|
||||
/// Evaluate several outputs at `x`.
|
||||
/// \complexity O(n) time. n is `depth`. One interior traversal per level from the memoizer resume index through the leaf. Extra space is the path memoizer (O(n) nodes, or one node if it does not memoize). A proof token adds one fold per level walked.
|
||||
template <std::size_t I0,
|
||||
std::size_t I1,
|
||||
std::size_t ...Is,
|
||||
|
|
@ -150,37 +235,6 @@ auto eval_point(const DpfKey & dpf, InputT && x, PathMemoizer && path = PathMemo
|
|||
*eval_point<Is>(dpf, x, path)...);
|
||||
}
|
||||
|
||||
/// Fold every point in `[from, to]` into `pi` (caller must `init_proof` first,
|
||||
/// or pass a fresh token via `prove_interval` below).
|
||||
template <typename KeyT, typename InputT>
|
||||
void prove_fold_interval(const KeyT & key, InputT from, InputT to,
|
||||
proof_token & pi)
|
||||
{
|
||||
static_assert(KeyT::is_verifiable,
|
||||
"prove_fold_interval: key must carry dpf::verifiable");
|
||||
using input_type = typename KeyT::input_type;
|
||||
auto cur = static_cast<input_type>(from);
|
||||
const auto last = static_cast<input_type>(to);
|
||||
for (;;)
|
||||
{
|
||||
nonmemoizing_path_memoizer<KeyT> path{};
|
||||
auto tx = key.offset_x(cur);
|
||||
utils::flip_msb_if_signed_integral(tx);
|
||||
internal::eval_point_interior(key, tx, path, &pi);
|
||||
if (cur == last)
|
||||
break;
|
||||
++cur;
|
||||
}
|
||||
}
|
||||
|
||||
/// Initialise `pr.token` and fold `[from, to]`.
|
||||
template <typename KeyT, typename InputT>
|
||||
void prove_interval(const KeyT & key, InputT from, InputT to, prove_ref pr)
|
||||
{
|
||||
detail::vdpf::init_proof(pr.token, key);
|
||||
prove_fold_interval(key, from, to, pr.token);
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_EVAL_POINT_HPP__
|
||||
|
|
|
|||
|
|
@ -24,26 +24,179 @@
|
|||
#include <utility>
|
||||
#include <tuple>
|
||||
#include <algorithm>
|
||||
#include <array>
|
||||
#include <iterator>
|
||||
#include <stdexcept>
|
||||
#include <list>
|
||||
#include <limits>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/dpf_key.hpp"
|
||||
#include "dpf/eval_common.hpp"
|
||||
#include "dpf/eval_target.hpp"
|
||||
#include "dpf/eval_point.hpp"
|
||||
#include "dpf/eval_interval.hpp"
|
||||
#include "dpf/path_memoizer.hpp"
|
||||
#include "dpf/sequence_memoizer.hpp"
|
||||
#include "dpf/sequence_utils.hpp"
|
||||
#include "dpf/subsequence_iterable.hpp"
|
||||
#include "dpf/subinterval_iterable.hpp"
|
||||
#include "dpf/verifiable.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
template <typename KeyT, typename ForwardIterator>
|
||||
void prove_sequence(const KeyT & key, ForwardIterator begin, ForwardIterator end,
|
||||
prove_ref pr);
|
||||
|
||||
// Contiguous runs at least this long use the interval tree (one expand per
|
||||
// node). Shorter runs stay on the path memoizer. Isolated points at least
|
||||
// this many use the breadth-first block walk, which beats a fresh path per
|
||||
// point once the list is wide.
|
||||
inline constexpr std::size_t sequence_interval_run = 24;
|
||||
// Breadth-first shares prefixes on sparse lists, but its per-level block-split
|
||||
// bookkeeping loses to a path memoizer once single-child expands cut the
|
||||
// path PRG cost in half. Keep the explicit `eval_sequence_breadth_first`
|
||||
// entry point; do not auto-select it from `eval_sequence`.
|
||||
inline constexpr std::size_t sequence_breadth_points = std::numeric_limits<std::size_t>::max();
|
||||
|
||||
template <std::size_t I,
|
||||
typename DpfKey,
|
||||
typename ForwardIterator,
|
||||
typename OutputBuffer>
|
||||
inline auto eval_sequence_breadth_first(const DpfKey & dpf, ForwardIterator begin,
|
||||
ForwardIterator end, OutputBuffer && outbuf);
|
||||
|
||||
template <bool Entire,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
typename InputT,
|
||||
typename OutputBuffers,
|
||||
std::size_t ...IIs>
|
||||
void scatter_interval_run(const DpfKey & dpf, InputT from, InputT to,
|
||||
OutputBuffers & outbufs, std::size_t index0, std::index_sequence<IIs...>)
|
||||
{
|
||||
// std::make_tuple, not utils::make_tuple: one output must stay a tuple
|
||||
// so each selected buffer is addressable by index.
|
||||
auto ibufs = std::make_tuple(
|
||||
make_output_buffer_for_interval<DpfKey, Is>(from, to)...);
|
||||
(void)eval_interval<Is...>(dpf, from, to, ibufs);
|
||||
|
||||
constexpr auto opl = DpfKey::outputs_per_leaf;
|
||||
constexpr auto lg = DpfKey::lg_outputs_per_leaf;
|
||||
constexpr auto mod_pow_2 = utils::mod_pow_2<InputT>{};
|
||||
const std::size_t preclip = mod_pow_2(dpf.offset_x(from), lg);
|
||||
|
||||
constexpr auto to_int = utils::to_integral_type<InputT>{};
|
||||
auto span = to_int(to) - to_int(from);
|
||||
constexpr auto bits = utils::bitlength_of_v<InputT>;
|
||||
if constexpr (bits < utils::bitlength_of_v<decltype(span)>)
|
||||
span &= (decltype(span){1} << bits) - 1;
|
||||
const std::size_t n = static_cast<std::size_t>(span) + 1;
|
||||
constexpr std::size_t leaf_mask = (std::size_t{1} << lg) - 1;
|
||||
|
||||
auto one = [&](auto which)
|
||||
{
|
||||
constexpr std::size_t k = decltype(which)::value;
|
||||
auto & src = utils::get<k>(ibufs);
|
||||
auto & dst = utils::get<k>(outbufs);
|
||||
auto put = [](auto & slot, const auto & val)
|
||||
{
|
||||
assign_share_slot(slot, val);
|
||||
};
|
||||
if constexpr (Entire)
|
||||
{
|
||||
for (std::size_t j = 0; j < n; ++j)
|
||||
{
|
||||
const std::size_t elem_i = preclip + j;
|
||||
const std::size_t leaf = elem_i - (elem_i & leaf_mask);
|
||||
for (std::size_t lane = 0; lane < opl; ++lane)
|
||||
put(dst[(index0 + j) * opl + lane], src[leaf + lane]);
|
||||
}
|
||||
}
|
||||
else
|
||||
{
|
||||
for (std::size_t j = 0; j < n; ++j)
|
||||
put(dst[index0 + j], src[preclip + j]);
|
||||
}
|
||||
};
|
||||
(one(std::integral_constant<std::size_t, IIs>{}), ...);
|
||||
}
|
||||
|
||||
namespace internal
|
||||
{
|
||||
|
||||
/// @brief Leaf interior nodes for a list, eight paths at a time.
|
||||
/// @details One scalar spine up to the shallowest divergence, then
|
||||
/// `tree::traverse8` across the whole suffix of each level. `walk[i]`
|
||||
/// is the offset input with the signed MSB already flipped, matching
|
||||
/// `eval_point`.
|
||||
template <typename DpfKey>
|
||||
std::vector<typename DpfKey::interior_node> sequence_wide_leaves(
|
||||
const DpfKey & dpf, const typename DpfKey::input_type * walk, std::size_t n)
|
||||
{
|
||||
using node = typename DpfKey::interior_node;
|
||||
std::vector<node> leaves(n);
|
||||
if (n == 0)
|
||||
return leaves;
|
||||
constexpr auto clz = utils::countl_zero_symmetric_difference<
|
||||
typename DpfKey::input_type>{};
|
||||
std::size_t common = dpf.depth + 1;
|
||||
for (std::size_t i = 1; i < n; ++i)
|
||||
common = std::min(common, clz(walk[0], walk[i]) + 1);
|
||||
|
||||
std::array<node, DpfKey::depth + 1> path{};
|
||||
path[0] = dpf.root();
|
||||
auto mask = dpf.msb_mask;
|
||||
for (std::size_t level = 1; level < common && level <= dpf.depth;
|
||||
++level, mask >>= 1)
|
||||
{
|
||||
const bool bit = (mask & walk[0]) != 0;
|
||||
path[level] = DpfKey::traverse_interior(path[level - 1],
|
||||
dpf.correction_word(level - 1, bit), bit,
|
||||
DpfKey::tree::is_last_level(level - 1, dpf.depth));
|
||||
}
|
||||
if (common > dpf.depth)
|
||||
{
|
||||
std::fill(leaves.begin(), leaves.end(), path[dpf.depth]);
|
||||
return leaves;
|
||||
}
|
||||
|
||||
std::vector<node> cur(n, path[common - 1]);
|
||||
std::vector<node> next(n);
|
||||
mask = dpf.msb_mask >> (common - 1);
|
||||
for (std::size_t level = common; level <= dpf.depth; ++level, mask >>= 1)
|
||||
{
|
||||
const node cw0 = dpf.correction_word(level - 1, false);
|
||||
const node cw1 = dpf.correction_word(level - 1, true);
|
||||
const bool last = DpfKey::tree::is_last_level(level - 1, dpf.depth);
|
||||
std::size_t i = 0;
|
||||
for (; i + 8 <= n; i += 8)
|
||||
{
|
||||
node parents[8];
|
||||
node cws[8];
|
||||
bool dirs[8];
|
||||
for (int k = 0; k < 8; ++k)
|
||||
{
|
||||
dirs[k] = (mask & walk[i + static_cast<std::size_t>(k)]) != 0;
|
||||
cws[k] = dirs[k] ? cw1 : cw0;
|
||||
parents[k] = cur[i + static_cast<std::size_t>(k)];
|
||||
}
|
||||
node outs[8];
|
||||
DpfKey::tree::traverse8(parents, cws, dirs, outs);
|
||||
for (int k = 0; k < 8; ++k)
|
||||
next[i + static_cast<std::size_t>(k)] = outs[k];
|
||||
}
|
||||
for (; i < n; ++i)
|
||||
{
|
||||
const bool bit = (mask & walk[i]) != 0;
|
||||
next[i] = DpfKey::traverse_interior(cur[i], bit ? cw1 : cw0, bit, last);
|
||||
}
|
||||
cur.swap(next);
|
||||
}
|
||||
return cur;
|
||||
}
|
||||
|
||||
template <std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
typename ForwardIterator,
|
||||
|
|
@ -53,24 +206,162 @@ auto eval_sequence_entire_node(const DpfKey & dpf, ForwardIterator begin, Forwar
|
|||
OutputBuffers && outbufs, std::index_sequence<IIs...>)
|
||||
{
|
||||
static constexpr std::size_t outputs_per_leaf = DpfKey::outputs_per_leaf;
|
||||
using input_type = typename DpfKey::input_type;
|
||||
constexpr bool any_packed =
|
||||
(utils::is_packed_subbyte_v<typename DpfKey::concrete_output_type<Is>> || ...);
|
||||
constexpr bool any_bit =
|
||||
(std::is_same_v<typename DpfKey::concrete_output_type<Is>, dpf::bit> || ...);
|
||||
constexpr bool integral_in = std::is_integral_v<input_type>;
|
||||
// Bit and packed leaves are proxy buffers. The interval scatter writes
|
||||
// ordinary word slots; those outputs stay on the path memoizer.
|
||||
constexpr bool interval_ok = integral_in && !any_packed && !any_bit;
|
||||
|
||||
auto path = make_basic_path_memoizer(dpf);
|
||||
|
||||
std::size_t i = 0;
|
||||
// DPF_UNROLL_LOOP
|
||||
for (auto it = begin; it != end; ++it, ++i)
|
||||
auto write_point = [&](auto & path, std::size_t i, const input_type & x)
|
||||
{
|
||||
if constexpr(utils::is_packed_subbyte_v<typename DpfKey::concrete_output_type<0>>)
|
||||
if constexpr (any_packed)
|
||||
{
|
||||
auto nodes = std::make_tuple(dpf::eval_point<Is>(dpf, *it, path).node...);
|
||||
auto nodes = std::make_tuple(dpf::eval_point<Is>(dpf, x, path).node...);
|
||||
(store_leaf_bytes(utils::get<IIs>(outbufs), i, std::get<IIs>(nodes)), ...);
|
||||
}
|
||||
else
|
||||
{
|
||||
auto temp = std::make_tuple(dpf::eval_point<Is>(dpf, *it, path).node...);
|
||||
(std::memcpy(&utils::get<IIs>(outbufs)[i*outputs_per_leaf], &utils::get<IIs>(temp), sizeof(typename DpfKey::concrete_output_type<Is>)*outputs_per_leaf), ...);
|
||||
auto temp = std::make_tuple(dpf::eval_point<Is>(dpf, x, path).node...);
|
||||
(utils::raw_memcpy(&utils::get<IIs>(outbufs)[i * outputs_per_leaf],
|
||||
&utils::get<IIs>(temp),
|
||||
sizeof(typename DpfKey::concrete_output_type<Is>) * outputs_per_leaf), ...);
|
||||
}
|
||||
};
|
||||
|
||||
if constexpr (interval_ok)
|
||||
{
|
||||
using iter_cat = typename std::iterator_traits<ForwardIterator>::iterator_category;
|
||||
bool classify = true;
|
||||
if constexpr (std::is_base_of_v<std::random_access_iterator_tag, iter_cat>)
|
||||
{
|
||||
// Below this length the interval-run cover does not apply, so skip
|
||||
// the classification pass.
|
||||
classify = static_cast<std::size_t>(std::distance(begin, end))
|
||||
>= sequence_interval_run;
|
||||
}
|
||||
if (classify)
|
||||
{
|
||||
bool sorted = true;
|
||||
std::size_t n = 0;
|
||||
std::size_t longest = 0;
|
||||
std::size_t cur = 0;
|
||||
bool dup = false;
|
||||
input_type prev{};
|
||||
for (auto it = begin; it != end; ++it, ++n)
|
||||
{
|
||||
if (n == 0)
|
||||
{
|
||||
prev = *it;
|
||||
longest = 1;
|
||||
cur = 1;
|
||||
continue;
|
||||
}
|
||||
const input_type x = *it;
|
||||
if (x < prev)
|
||||
sorted = false;
|
||||
else if (x == prev)
|
||||
dup = true;
|
||||
if (x == static_cast<input_type>(prev + input_type{1}))
|
||||
{
|
||||
++cur;
|
||||
if (cur > longest)
|
||||
longest = cur;
|
||||
}
|
||||
else
|
||||
cur = 1;
|
||||
prev = x;
|
||||
}
|
||||
if (sorted && longest >= sequence_interval_run)
|
||||
{
|
||||
std::size_t i = 0;
|
||||
for (auto it = begin; it != end; )
|
||||
{
|
||||
const input_type run_from = *it;
|
||||
input_type run_to = run_from;
|
||||
auto run_end = it;
|
||||
++run_end;
|
||||
while (run_end != end
|
||||
&& *run_end == static_cast<input_type>(run_to + input_type{1}))
|
||||
{
|
||||
run_to = *run_end;
|
||||
++run_end;
|
||||
}
|
||||
const auto span_n = static_cast<std::size_t>(std::distance(it, run_end));
|
||||
if (span_n >= sequence_interval_run)
|
||||
{
|
||||
scatter_interval_run<true, Is...>(dpf, run_from, run_to, outbufs, i,
|
||||
std::index_sequence<IIs...>{});
|
||||
i += span_n;
|
||||
}
|
||||
else
|
||||
{
|
||||
auto path = make_basic_path_memoizer(dpf);
|
||||
for (auto p = it; p != run_end; ++p, ++i)
|
||||
write_point(path, i, *p);
|
||||
}
|
||||
it = run_end;
|
||||
}
|
||||
return utils::make_tuple(
|
||||
dpf::subsequence_iterable<DpfKey, decltype(std::begin(utils::get<IIs>(outbufs))), ForwardIterator>(std::begin(utils::get<IIs>(outbufs)), begin, end)...);
|
||||
}
|
||||
// Breadth-first pays for a block split at every level. On a short
|
||||
// tree (8-bit inputs) that overhead loses to the path memoizer.
|
||||
// The one-output overload is the only one; keep it out of the
|
||||
// multi-output instantiation.
|
||||
if constexpr (sizeof...(Is) == 1)
|
||||
{
|
||||
if (sorted && !dup && n >= sequence_breadth_points
|
||||
&& DpfKey::depth >= 16)
|
||||
{
|
||||
eval_sequence_breadth_first<Is...>(dpf, begin, end, utils::get<0>(outbufs));
|
||||
return utils::make_tuple(
|
||||
dpf::subsequence_iterable<DpfKey, decltype(std::begin(utils::get<IIs>(outbufs))), ForwardIterator>(std::begin(utils::get<IIs>(outbufs)), begin, end)...);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const std::size_t nseq = static_cast<std::size_t>(std::distance(begin, end));
|
||||
if constexpr (sizeof...(Is) == 1 && !DpfKey::is_verifiable
|
||||
&& prg_has_indep4<typename DpfKey::interior_prg>::value)
|
||||
{
|
||||
if (nseq >= 16)
|
||||
{
|
||||
using input_type = typename DpfKey::input_type;
|
||||
std::vector<input_type> walk(nseq);
|
||||
std::vector<input_type> lane(nseq);
|
||||
std::size_t i = 0;
|
||||
for (auto it = begin; it != end; ++it, ++i)
|
||||
{
|
||||
lane[i] = dpf.offset_x(*it);
|
||||
walk[i] = lane[i];
|
||||
utils::flip_msb_if_signed_integral(walk[i]);
|
||||
}
|
||||
const auto leaves = sequence_wide_leaves(dpf, walk.data(), nseq);
|
||||
constexpr std::size_t ids[] = {Is...};
|
||||
constexpr std::size_t I0 = ids[0];
|
||||
using output_type = typename DpfKey::concrete_output_type<I0>;
|
||||
auto & dst = utils::get<0>(outbufs);
|
||||
for (std::size_t p = 0; p < nseq; ++p)
|
||||
{
|
||||
auto wrapped = make_eval_dpf_output<DpfKey, output_type>(
|
||||
dpf.template traverse_exterior<I0>(leaves[p]), lane[p]);
|
||||
utils::raw_memcpy(&dst[p * outputs_per_leaf], &wrapped.node,
|
||||
sizeof(output_type) * outputs_per_leaf);
|
||||
}
|
||||
return utils::make_tuple(
|
||||
dpf::subsequence_iterable<DpfKey, decltype(std::begin(utils::get<IIs>(outbufs))), ForwardIterator>(std::begin(utils::get<IIs>(outbufs)), begin, end)...);
|
||||
}
|
||||
}
|
||||
|
||||
auto path = make_basic_path_memoizer(dpf);
|
||||
std::size_t i = 0;
|
||||
for (auto it = begin; it != end; ++it, ++i)
|
||||
write_point(path, i, *it);
|
||||
return utils::make_tuple(
|
||||
dpf::subsequence_iterable<DpfKey, decltype(std::begin(utils::get<IIs>(outbufs))), ForwardIterator>(std::begin(utils::get<IIs>(outbufs)), begin, end)...);
|
||||
}
|
||||
|
|
@ -79,23 +370,7 @@ template <typename Slot, typename Val>
|
|||
HEDLEY_ALWAYS_INLINE
|
||||
void assign_eval_slot(Slot && slot, Val && val)
|
||||
{
|
||||
using val_t = std::decay_t<Val>;
|
||||
if constexpr (is_secret_share_v<std::decay_t<Slot>>)
|
||||
{
|
||||
using elem_t = std::decay_t<Slot>;
|
||||
if constexpr (is_secret_share_v<val_t>)
|
||||
slot = elem_t::from_raw(val.raw());
|
||||
else
|
||||
slot = elem_t::from_raw(static_cast<typename elem_t::value_type>(val));
|
||||
}
|
||||
else if constexpr (is_secret_share_v<val_t>)
|
||||
{
|
||||
slot = val.raw();
|
||||
}
|
||||
else
|
||||
{
|
||||
slot = std::forward<Val>(val);
|
||||
}
|
||||
assign_share_slot(slot, std::forward<Val>(val));
|
||||
}
|
||||
|
||||
template <std::size_t ...Is,
|
||||
|
|
@ -106,24 +381,101 @@ template <std::size_t ...Is,
|
|||
auto eval_sequence_output_only(const DpfKey & dpf, ForwardIterator begin, ForwardIterator end, OutputBuffers && outbufs,
|
||||
std::index_sequence<IIs...>)
|
||||
{
|
||||
auto path = make_basic_path_memoizer(dpf);
|
||||
using input_type = typename DpfKey::input_type;
|
||||
constexpr bool any_packed =
|
||||
(utils::is_packed_subbyte_v<typename DpfKey::concrete_output_type<Is>> || ...);
|
||||
constexpr bool any_bit =
|
||||
(std::is_same_v<typename DpfKey::concrete_output_type<Is>, dpf::bit> || ...);
|
||||
constexpr bool integral_in = std::is_integral_v<input_type>;
|
||||
|
||||
std::size_t i = 0;
|
||||
// DPF_UNROLL_LOOP
|
||||
for (auto it = begin; it != end; ++it, ++i)
|
||||
auto write_point = [&](auto & path, std::size_t i, const input_type & x)
|
||||
{
|
||||
(assign_eval_slot(utils::get<IIs>(outbufs)[i],
|
||||
*dpf::eval_point<Is>(dpf, *it, path)), ...);
|
||||
}
|
||||
if (i == 0)
|
||||
*dpf::eval_point<Is>(dpf, x, path)), ...);
|
||||
};
|
||||
|
||||
auto finish = [&](std::size_t i)
|
||||
{
|
||||
return utils::make_tuple(subinterval_iterable(std::begin(utils::get<IIs>(outbufs)), utils::size(utils::get<IIs>(outbufs)), 0, 0, 0, 0, false)...);
|
||||
if (i == 0)
|
||||
{
|
||||
return utils::make_tuple(subinterval_iterable(std::begin(utils::get<IIs>(outbufs)), utils::size(utils::get<IIs>(outbufs)), 0, 0, 0, 0, false)...);
|
||||
}
|
||||
return utils::make_tuple(subinterval_iterable(std::begin(utils::get<IIs>(outbufs)), utils::size(utils::get<IIs>(outbufs)), 0, i - 1, 0, 0)...);
|
||||
};
|
||||
|
||||
if constexpr (integral_in && !any_packed && !any_bit)
|
||||
{
|
||||
bool sorted = true;
|
||||
std::size_t longest = 0;
|
||||
std::size_t cur = 0;
|
||||
std::size_t n = 0;
|
||||
input_type prev{};
|
||||
for (auto it = begin; it != end; ++it, ++n)
|
||||
{
|
||||
if (n == 0)
|
||||
{
|
||||
prev = *it;
|
||||
longest = 1;
|
||||
cur = 1;
|
||||
continue;
|
||||
}
|
||||
const input_type x = *it;
|
||||
if (x < prev)
|
||||
sorted = false;
|
||||
if (x == static_cast<input_type>(prev + input_type{1}))
|
||||
{
|
||||
++cur;
|
||||
if (cur > longest)
|
||||
longest = cur;
|
||||
}
|
||||
else
|
||||
cur = 1;
|
||||
prev = x;
|
||||
}
|
||||
if (sorted && longest >= sequence_interval_run)
|
||||
{
|
||||
auto path = make_basic_path_memoizer(dpf);
|
||||
std::size_t i = 0;
|
||||
for (auto it = begin; it != end; )
|
||||
{
|
||||
const input_type run_from = *it;
|
||||
input_type run_to = run_from;
|
||||
auto run_end = it;
|
||||
++run_end;
|
||||
while (run_end != end
|
||||
&& *run_end == static_cast<input_type>(run_to + input_type{1}))
|
||||
{
|
||||
run_to = *run_end;
|
||||
++run_end;
|
||||
}
|
||||
const auto span_n = static_cast<std::size_t>(std::distance(it, run_end));
|
||||
if (span_n >= sequence_interval_run)
|
||||
{
|
||||
scatter_interval_run<false, Is...>(dpf, run_from, run_to, outbufs, i,
|
||||
std::index_sequence<IIs...>{});
|
||||
i += span_n;
|
||||
}
|
||||
else
|
||||
{
|
||||
for (auto p = it; p != run_end; ++p, ++i)
|
||||
write_point(path, i, *p);
|
||||
}
|
||||
it = run_end;
|
||||
}
|
||||
return finish(i);
|
||||
}
|
||||
}
|
||||
return utils::make_tuple(subinterval_iterable(std::begin(utils::get<IIs>(outbufs)), utils::size(utils::get<IIs>(outbufs)), 0, i-1, 0, 0)...);
|
||||
|
||||
auto path = make_basic_path_memoizer(dpf);
|
||||
std::size_t i = 0;
|
||||
for (auto it = begin; it != end; ++it, ++i)
|
||||
write_point(path, i, *it);
|
||||
return finish(i);
|
||||
}
|
||||
|
||||
} // namespace internal
|
||||
|
||||
/// \complexity O(n k) interior traversals in the worst case and O(k) node workspace. k is the number of listed points and n is `depth`. The breadth-first buffer is 2k nodes, so each level traverses at most one node per point. Shared prefixes do fewer traversals. A recipe memoizer instead stores O(recipe leaf nodes) (see that memoizer).
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
|
|
@ -136,6 +488,7 @@ template <std::size_t I = 0,
|
|||
inline auto eval_sequence(const DpfKey & dpf, ForwardIterator begin, ForwardIterator end,
|
||||
OutputBuffers && outbufs, ReturnType return_type = ReturnType{})
|
||||
{
|
||||
(void)return_type;
|
||||
static_assert(std::is_same_v<ReturnType, return_entire_node_tag_> ||
|
||||
std::is_same_v<ReturnType, return_output_only_tag_>);
|
||||
if constexpr(std::is_same_v<ReturnType, return_entire_node_tag_>)
|
||||
|
|
@ -148,6 +501,54 @@ inline auto eval_sequence(const DpfKey & dpf, ForwardIterator begin, ForwardIter
|
|||
}
|
||||
}
|
||||
|
||||
/// @brief Evaluate a sorted sequence and fold the same VDPF proof as `prove_sequence`.
|
||||
/// @details Uses interval-run covers (not a path-memo double walk). Values are
|
||||
/// written by a single subsequent `eval_sequence` without proof.
|
||||
/// \complexity O(n k) interior traversals in the worst case and O(k) node workspace. k is the number of listed points and n is `depth`. The breadth-first buffer is 2k nodes, so each level traverses at most one node per point. Shared prefixes do fewer traversals. A recipe memoizer instead stores O(recipe leaf nodes) (see that memoizer).
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
typename ForwardIterator,
|
||||
typename OutputBuffers,
|
||||
typename ReturnType = return_entire_node_tag_,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey> && !is_multilevel_key_v<DpfKey>, bool> = true,
|
||||
std::enable_if_t<std::is_base_of_v<return_type_tag_, ReturnType>, bool> = true,
|
||||
std::enable_if_t<!std::is_base_of_v<return_type_tag_, OutputBuffers>, bool> = true>
|
||||
inline auto eval_sequence(const DpfKey & dpf, ForwardIterator begin, ForwardIterator end,
|
||||
OutputBuffers && outbufs, prove_ref pr, ReturnType return_type = ReturnType{})
|
||||
{
|
||||
static_assert(DpfKey::is_verifiable,
|
||||
"eval_sequence(..., prove(π)): key must carry dpf::verifiable");
|
||||
prove_sequence(dpf, begin, end, pr);
|
||||
return eval_sequence<I, Is...>(dpf, begin, end,
|
||||
std::forward<OutputBuffers>(outbufs), return_type);
|
||||
}
|
||||
|
||||
/// @brief Evaluate a sorted sequence and fold each written output into a sketch.
|
||||
/// \complexity O(n k) interior traversals in the worst case and O(k) node workspace. k is the number of listed points and n is `depth`. The breadth-first buffer is 2k nodes, so each level traverses at most one node per point. Shared prefixes do fewer traversals. A recipe memoizer instead stores O(recipe leaf nodes) (see that memoizer).
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
typename ForwardIterator,
|
||||
typename OutputBuffers,
|
||||
typename ReturnType = return_entire_node_tag_,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey> && !is_multilevel_key_v<DpfKey>, bool> = true,
|
||||
std::enable_if_t<std::is_base_of_v<return_type_tag_, ReturnType>, bool> = true,
|
||||
std::enable_if_t<!std::is_base_of_v<return_type_tag_, OutputBuffers>, bool> = true>
|
||||
inline auto eval_sequence(const DpfKey & dpf, ForwardIterator begin, ForwardIterator end,
|
||||
OutputBuffers && outbufs, sketch_ref & sk, ReturnType return_type = ReturnType{})
|
||||
{
|
||||
static_assert(DpfKey::is_extractable,
|
||||
"eval_sequence(..., sketch(σ)): key must carry dpf::extractable");
|
||||
auto ret = eval_sequence<I, Is...>(dpf, begin, end, outbufs, return_type);
|
||||
if constexpr (sizeof...(Is) == 0)
|
||||
{
|
||||
for (std::size_t k = 0; k < utils::size(outbufs); ++k)
|
||||
sk.absorb(outbufs[k]);
|
||||
}
|
||||
return ret;
|
||||
}
|
||||
|
||||
/// @brief Evaluate the sorted range `[begin, end)`, allocating a buffer.
|
||||
/// @tparam I output index
|
||||
/// @tparam Is is
|
||||
|
|
@ -161,6 +562,7 @@ inline auto eval_sequence(const DpfKey & dpf, ForwardIterator begin, ForwardIter
|
|||
/// @param begin the iterator to the first query
|
||||
/// @param end the iterator past the last query
|
||||
/// @return Pair of buffer (or tuple of buffers) and an iterable in list order.
|
||||
/// \complexity O(n k) interior traversals in the worst case and O(k) node workspace. k is the number of listed points and n is `depth`. The breadth-first buffer is 2k nodes, so each level traverses at most one node per point. Shared prefixes do fewer traversals. A recipe memoizer instead stores O(recipe leaf nodes) (see that memoizer).
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
|
|
@ -182,6 +584,7 @@ auto eval_sequence(const DpfKey & dpf, ForwardIterator begin, ForwardIterator en
|
|||
return std::make_pair(std::move(outbufs), std::move(iterable));
|
||||
}
|
||||
|
||||
/// \complexity O(n k) interior traversals in the worst case and O(k) node workspace. k is the number of listed points and n is `depth`. The breadth-first buffer is 2k nodes, so each level traverses at most one node per point. Shared prefixes do fewer traversals. A recipe memoizer instead stores O(recipe leaf nodes) (see that memoizer).
|
||||
template <std::size_t I = 0,
|
||||
typename DpfKey,
|
||||
typename ForwardIterator,
|
||||
|
|
@ -219,7 +622,10 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
bool curhalf = (dpf_type::depth ^ 1) & 1;
|
||||
memo[!curhalf*nodes_in_sequence + 0] = dpf.root();
|
||||
|
||||
std::list<ForwardIterator> splits{begin, end};
|
||||
std::vector<ForwardIterator> splits;
|
||||
splits.reserve(nodes_in_sequence + 1);
|
||||
splits.push_back(begin);
|
||||
splits.push_back(end);
|
||||
|
||||
std::size_t level_index = 1;
|
||||
auto func = [&](const bool flip = false)
|
||||
|
|
@ -233,16 +639,18 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
dpf_type::depth);
|
||||
// `lower` and `upper` are always adjacent elements of `splits` with `lower` < `upper`
|
||||
// [lower, upper) = "block"
|
||||
for (auto upper = std::begin(splits), lower = upper++; upper != std::end(splits); lower = upper++)
|
||||
for (std::size_t s = 0; s + 1 < splits.size(); ++s)
|
||||
{
|
||||
auto lower = splits[s];
|
||||
auto upper = splits[s + 1];
|
||||
// `upper_bound()` returns iterator to first element where the relevant bit (based on `mask`) is set
|
||||
auto it = std::upper_bound(*lower, *upper, mask,
|
||||
auto it = std::upper_bound(lower, upper, mask,
|
||||
[&flip](auto a, auto b){ return static_cast<bool>(a&b) ^ flip; });
|
||||
if (it == *lower) // right only since first element in "block" requires right traversal
|
||||
if (it == lower) // right only since first element in "block" requires right traversal
|
||||
{
|
||||
memo[curhalf*nodes_in_sequence + i++] = dpf_type::traverse_interior(memo[!curhalf*nodes_in_sequence + j++], cw[1], 1, is_last);
|
||||
}
|
||||
else if (it == *upper) // left only since no element in "block" requires right traversal
|
||||
else if (it == upper) // left only since no element in "block" requires right traversal
|
||||
{
|
||||
memo[curhalf*nodes_in_sequence + i++] = dpf_type::traverse_interior(memo[!curhalf*nodes_in_sequence + j++], cw[0], 0, is_last);
|
||||
}
|
||||
|
|
@ -252,7 +660,8 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
auto kids = dpf_type::traverse_interior01(cur_node, cw[0], cw[1], is_last);
|
||||
memo[curhalf*nodes_in_sequence + i++] = kids[0];
|
||||
memo[curhalf*nodes_in_sequence + i++] = kids[1];
|
||||
splits.insert(upper, it);
|
||||
splits.insert(splits.begin() + static_cast<std::ptrdiff_t>(s + 1), it);
|
||||
++s; // skip the newly inserted right-block start
|
||||
}
|
||||
}
|
||||
};
|
||||
|
|
@ -287,7 +696,7 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
}
|
||||
else
|
||||
{
|
||||
std::memcpy(&outbuf[i*dpf_type::outputs_per_leaf], &leaf, sizeof(output_type)*dpf_type::outputs_per_leaf);
|
||||
utils::raw_memcpy(&outbuf[i*dpf_type::outputs_per_leaf], &leaf, sizeof(output_type)*dpf_type::outputs_per_leaf);
|
||||
}
|
||||
prev = curr++;
|
||||
}
|
||||
|
|
@ -295,6 +704,7 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
return subsequence_iterable<DpfKey, decltype(std::begin(outbuf)), ForwardIterator>(std::begin(outbuf), begin, end);
|
||||
}
|
||||
|
||||
/// \complexity O(n k) interior traversals in the worst case and O(k) node workspace. k is the number of listed points and n is `depth`. The breadth-first buffer is 2k nodes, so each level traverses at most one node per point. Shared prefixes do fewer traversals. A recipe memoizer instead stores O(recipe leaf nodes) (see that memoizer).
|
||||
template <std::size_t I = 0,
|
||||
typename DpfKey,
|
||||
typename ForwardIterator>
|
||||
|
|
@ -387,7 +797,7 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
}
|
||||
else
|
||||
{
|
||||
std::memcpy(&outbuf[j*dpf_type::outputs_per_leaf], &leaf, sizeof(output_type)*dpf_type::outputs_per_leaf);
|
||||
utils::raw_memcpy(&outbuf[j*dpf_type::outputs_per_leaf], &leaf, sizeof(output_type)*dpf_type::outputs_per_leaf);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -426,11 +836,7 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
}
|
||||
auto v = extract_leaf<node_type, output_type>(node,
|
||||
recipe.output_indices()[i] % dpf_type::outputs_per_leaf);
|
||||
using elem_t = std::decay_t<decltype(outbuf[i])>;
|
||||
if constexpr (is_secret_share_v<elem_t>)
|
||||
outbuf[i] = elem_t::from_raw(v);
|
||||
else
|
||||
outbuf[i] = v;
|
||||
assign_share_slot(outbuf[i], v);
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -444,6 +850,7 @@ template <std::size_t ...Is,
|
|||
auto eval_sequence(const DpfKey & dpf, const sequence_recipe & recipe,
|
||||
OutputBuffers && outbufs, SequenceMemoizer && memoizer, ReturnType return_type, std::index_sequence<IIs...>)
|
||||
{
|
||||
(void)return_type;
|
||||
internal::eval_sequence_interior(dpf, recipe, memoizer);
|
||||
|
||||
static_assert(std::is_same_v<ReturnType, return_entire_node_tag_> ||
|
||||
|
|
@ -483,6 +890,7 @@ auto eval_sequence(const DpfKey & dpf, const sequence_recipe & recipe,
|
|||
/// @param memoizer the memoizer built for this key
|
||||
/// @param return_type `return_entire_node_tag_{}` or `return_output_only_tag_{}`
|
||||
/// @return the evaluation result
|
||||
/// \complexity O(n k) interior traversals in the worst case and O(k) node workspace. k is the number of listed points and n is `depth`. The breadth-first buffer is 2k nodes, so each level traverses at most one node per point. Shared prefixes do fewer traversals. A recipe memoizer instead stores O(recipe leaf nodes) (see that memoizer).
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
|
|
@ -502,6 +910,7 @@ auto eval_sequence(const DpfKey & dpf, const sequence_recipe & recipe,
|
|||
return internal::eval_sequence<I, Is...>(dpf, recipe, outbufs, memoizer, return_type, std::make_index_sequence<1+sizeof...(Is)>());
|
||||
}
|
||||
|
||||
/// \complexity O(n k) interior traversals in the worst case and O(k) node workspace. k is the number of listed points and n is `depth`. The breadth-first buffer is 2k nodes, so each level traverses at most one node per point. Shared prefixes do fewer traversals. A recipe memoizer instead stores O(recipe leaf nodes) (see that memoizer).
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
|
|
@ -518,6 +927,7 @@ auto eval_sequence(const DpfKey & dpf, const sequence_recipe & recipe,
|
|||
dpf::make_double_space_sequence_memoizer<DpfKey>(recipe), return_type);
|
||||
}
|
||||
|
||||
/// \complexity O(n k) interior traversals in the worst case and O(k) node workspace. k is the number of listed points and n is `depth`. The breadth-first buffer is 2k nodes, so each level traverses at most one node per point. Shared prefixes do fewer traversals. A recipe memoizer instead stores O(recipe leaf nodes) (see that memoizer).
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
|
|
@ -541,6 +951,7 @@ auto eval_sequence(const DpfKey & dpf, const sequence_recipe & recipe,
|
|||
return std::make_pair(std::move(outbufs), std::move(iterable));
|
||||
}
|
||||
|
||||
/// \complexity O(n k) interior traversals in the worst case and O(k) node workspace. k is the number of listed points and n is `depth`. The breadth-first buffer is 2k nodes, so each level traverses at most one node per point. Shared prefixes do fewer traversals. A recipe memoizer instead stores O(recipe leaf nodes) (see that memoizer).
|
||||
template <std::size_t I = 0,
|
||||
std::size_t ...Is,
|
||||
typename DpfKey,
|
||||
|
|
@ -554,6 +965,38 @@ auto eval_sequence(const DpfKey & dpf, const sequence_recipe & recipe,
|
|||
dpf::make_double_space_sequence_memoizer<DpfKey>(recipe), return_type);
|
||||
}
|
||||
|
||||
/// @brief Fold a sorted sequence into `pr` via interval-run covers.
|
||||
/// @details Maximal contiguous runs use once-per-BFS-node `prove_fold_interval`.
|
||||
/// Isolated points are length-1 runs. Both parties must see the same
|
||||
/// sorted list. Prefer this over path-memo when the list is dense.
|
||||
template <typename KeyT, typename ForwardIterator>
|
||||
void prove_sequence(const KeyT & key, ForwardIterator begin, ForwardIterator end,
|
||||
prove_ref pr)
|
||||
{
|
||||
static_assert(KeyT::is_verifiable,
|
||||
"prove_sequence: key must carry dpf::verifiable");
|
||||
if (HEDLEY_UNLIKELY(begin != end && !std::is_sorted(begin, end)))
|
||||
throw std::runtime_error("list must be sorted");
|
||||
detail::vdpf::init_proof(pr.token, key);
|
||||
using input_type = typename KeyT::input_type;
|
||||
for (auto it = begin; it != end; )
|
||||
{
|
||||
const auto run_from = static_cast<input_type>(*it);
|
||||
auto run_to = run_from;
|
||||
++it;
|
||||
while (it != end)
|
||||
{
|
||||
const auto next = static_cast<input_type>(*it);
|
||||
if (next != static_cast<input_type>(run_to + input_type{1}))
|
||||
break;
|
||||
run_to = next;
|
||||
++it;
|
||||
}
|
||||
prove_fold_interval(key, run_from, run_to, pr.token);
|
||||
}
|
||||
detail::vdpf::fold_output_binding(pr.token, key);
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_EVAL_SEQUENCE_HPP__
|
||||
|
|
|
|||
|
|
@ -135,6 +135,66 @@ template <typename T>
|
|||
inline constexpr bool is_multilevel_key_v =
|
||||
is_multilevel_key<std::decay_t<T>>::value;
|
||||
|
||||
/// @brief True when `T` stores its DPF key in a public member named `dpf_key`
|
||||
/// and is not itself a DPF key.
|
||||
/// @details Wrappers (`ic_key`, `vdpf_plus_key`, `dpf3_cmp_key`, `dpf3_ic_key`,
|
||||
/// and the distributed `opened_*` envelopes) use that member name.
|
||||
/// `party_key` inherits the DPF key, so it is not a wrapper.
|
||||
template <typename T, typename = void>
|
||||
struct has_embedded_dpf_key : std::false_type
|
||||
{
|
||||
};
|
||||
template <typename T>
|
||||
struct has_embedded_dpf_key<T,
|
||||
std::void_t<decltype(std::declval<T &>().dpf_key)>>
|
||||
: std::bool_constant<!looks_like_dpf_key_v<T>>
|
||||
{
|
||||
};
|
||||
template <typename T>
|
||||
inline constexpr bool has_embedded_dpf_key_v =
|
||||
has_embedded_dpf_key<std::decay_t<T>>::value;
|
||||
|
||||
/// @brief The DPF key `T` evaluates as: `T` itself, or `T::dpf_key` peeled
|
||||
/// until the result looks like a DPF key.
|
||||
template <typename T, typename = void>
|
||||
struct bare_dpf_key
|
||||
{
|
||||
using type = std::decay_t<T>;
|
||||
};
|
||||
template <typename T>
|
||||
struct bare_dpf_key<T, std::enable_if_t<has_embedded_dpf_key_v<std::decay_t<T>>>>
|
||||
: bare_dpf_key<std::decay_t<decltype(std::declval<std::decay_t<T> &>().dpf_key)>>
|
||||
{
|
||||
};
|
||||
template <typename T>
|
||||
using bare_dpf_key_t = typename bare_dpf_key<std::decay_t<T>>::type;
|
||||
|
||||
/// @brief Key-first protocol evals (`dpf3`, `dpf3_cmp`, `dpf3_ic`) keep their
|
||||
/// own overloads. Peel does not replace those.
|
||||
template <typename T, typename = void>
|
||||
struct flag_is_dpf3 : std::false_type {};
|
||||
template <typename T>
|
||||
struct flag_is_dpf3<T, std::void_t<decltype(T::is_dpf3)>>
|
||||
: std::bool_constant<T::is_dpf3> {};
|
||||
|
||||
template <typename T, typename = void>
|
||||
struct flag_is_dpf3_cmp : std::false_type {};
|
||||
template <typename T>
|
||||
struct flag_is_dpf3_cmp<T, std::void_t<decltype(T::is_dpf3_cmp)>>
|
||||
: std::bool_constant<T::is_dpf3_cmp> {};
|
||||
|
||||
template <typename T, typename = void>
|
||||
struct flag_is_dpf3_ic : std::false_type {};
|
||||
template <typename T>
|
||||
struct flag_is_dpf3_ic<T, std::void_t<decltype(T::is_dpf3_ic)>>
|
||||
: std::bool_constant<T::is_dpf3_ic> {};
|
||||
|
||||
template <typename T>
|
||||
inline constexpr bool owns_protocol_eval_v =
|
||||
flag_is_dpf3<std::decay_t<T>>::value
|
||||
|| flag_is_dpf3_cmp<std::decay_t<T>>::value
|
||||
|| flag_is_dpf3_ic<std::decay_t<T>>::value;
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_EVAL_TARGET_HPP__
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@
|
|||
#include <cstring>
|
||||
#include <algorithm>
|
||||
#include <iterator>
|
||||
#include <limits>
|
||||
#include <list>
|
||||
#include <stdexcept>
|
||||
#include <type_traits>
|
||||
|
|
@ -36,6 +37,7 @@
|
|||
#include "dpf/interval_memoizer.hpp"
|
||||
#include "dpf/aligned_allocator.hpp"
|
||||
#include "dpf/leaf_node.hpp"
|
||||
#include "dpf/verifiable.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
|
@ -70,8 +72,10 @@ constexpr std::size_t resolved_out_prefix() noexcept
|
|||
// eval_point(target, key, x [, path])
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// \complexity O(n) time. n is `depth`. One interior traversal per level from the memoizer resume index through the leaf. Extra space is the path memoizer (O(n) nodes, or one node if it does not memoize). A proof token adds one fold per level walked.
|
||||
template <std::size_t I, std::size_t N, typename KeyT, typename QueryT,
|
||||
typename PathMemoizer = nonmemoizing_path_memoizer<KeyT>>
|
||||
typename PathMemoizer = nonmemoizing_path_memoizer<KeyT>,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_point(out_t<I, N>, const KeyT & key, QueryT && x,
|
||||
PathMemoizer && path = PathMemoizer{})
|
||||
{
|
||||
|
|
@ -88,8 +92,10 @@ auto eval_point(out_t<I, N>, const KeyT & key, QueryT && x,
|
|||
}
|
||||
}
|
||||
|
||||
/// \complexity O(n) time. n is `depth`. One interior traversal per level from the memoizer resume index through the leaf. Extra space is the path memoizer (O(n) nodes, or one node if it does not memoize). A proof token adds one fold per level walked.
|
||||
template <typename Beta = uint64_t, typename KeyT, typename QueryT,
|
||||
typename PathMemoizer = basic_path_memoizer<KeyT>>
|
||||
typename PathMemoizer = basic_path_memoizer<KeyT>,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_point(cmp_t, const KeyT & key, QueryT && x,
|
||||
PathMemoizer && path = PathMemoizer{})
|
||||
{
|
||||
|
|
@ -97,8 +103,41 @@ auto eval_point(cmp_t, const KeyT & key, QueryT && x,
|
|||
std::forward<PathMemoizer>(path));
|
||||
}
|
||||
|
||||
/// \complexity O(n) time. n is `depth`. One interior traversal per level from the memoizer resume index through the leaf. Extra space is the path memoizer (O(n) nodes, or one node if it does not memoize). A proof token adds one fold per level walked.
|
||||
template <typename Beta = uint64_t, typename KeyT, typename QueryT,
|
||||
typename PathMemoizer = basic_path_memoizer<KeyT>,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_point(cmp_t, const KeyT & key, QueryT && x, prove_ref pr,
|
||||
PathMemoizer && path = PathMemoizer{})
|
||||
{
|
||||
static_assert(KeyT::is_verifiable,
|
||||
"eval_point(cmp, ..., prove(π)): key must carry dpf::verifiable");
|
||||
detail::vdpf::init_proof(pr.token, key);
|
||||
auto out = detail::incr::eval_cmp_point_impl<Beta>(key, std::forward<QueryT>(x),
|
||||
std::forward<PathMemoizer>(path), &pr.token);
|
||||
detail::vdpf::fold_output_binding(pr.token, key);
|
||||
return out;
|
||||
}
|
||||
|
||||
/// \complexity O(n) time. n is `depth`. One interior traversal per level from the memoizer resume index through the leaf. Extra space is the path memoizer (O(n) nodes, or one node if it does not memoize). A proof token adds one fold per level walked.
|
||||
template <typename Beta = uint64_t, typename KeyT, typename QueryT,
|
||||
typename PathMemoizer = basic_path_memoizer<KeyT>,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_point(cmp_t, const KeyT & key, QueryT && x, sketch_ref & sk,
|
||||
PathMemoizer && path = PathMemoizer{})
|
||||
{
|
||||
static_assert(KeyT::is_extractable,
|
||||
"eval_point(cmp, ..., sketch(σ)): key must carry dpf::extractable");
|
||||
auto y = detail::incr::eval_cmp_point_impl<Beta>(key, std::forward<QueryT>(x),
|
||||
std::forward<PathMemoizer>(path));
|
||||
sk.absorb(y);
|
||||
return y;
|
||||
}
|
||||
|
||||
/// \complexity O(n) time. n is `depth`. One interior traversal per level from the memoizer resume index through the leaf. Extra space is the path memoizer (O(n) nodes, or one node if it does not memoize). A proof token adds one fold per level walked.
|
||||
template <std::size_t L, typename Beta = uint64_t, typename KeyT, typename QueryT,
|
||||
typename PathMemoizer = basic_path_memoizer<KeyT>>
|
||||
typename PathMemoizer = basic_path_memoizer<KeyT>,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_point(cmp_prefix_t<L>, const KeyT & key, QueryT && x,
|
||||
PathMemoizer && path = PathMemoizer{})
|
||||
{
|
||||
|
|
@ -106,12 +145,30 @@ auto eval_point(cmp_prefix_t<L>, const KeyT & key, QueryT && x,
|
|||
std::forward<QueryT>(x), std::forward<PathMemoizer>(path));
|
||||
}
|
||||
|
||||
/// \complexity O(n) time. n is `depth`. One interior traversal per level from the memoizer resume index through the leaf. Extra space is the path memoizer (O(n) nodes, or one node if it does not memoize). A proof token adds one fold per level walked.
|
||||
template <std::size_t L, typename Beta = uint64_t, typename KeyT, typename QueryT,
|
||||
typename PathMemoizer = basic_path_memoizer<KeyT>,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_point(cmp_prefix_t<L>, const KeyT & key, QueryT && x, prove_ref pr,
|
||||
PathMemoizer && path = PathMemoizer{})
|
||||
{
|
||||
static_assert(KeyT::is_verifiable,
|
||||
"eval_point(cmp_prefix, ..., prove(π)): key must carry dpf::verifiable");
|
||||
detail::vdpf::init_proof(pr.token, key);
|
||||
auto out = detail::incr::eval_cmp_prefix_point_impl<L, Beta>(key,
|
||||
std::forward<QueryT>(x), std::forward<PathMemoizer>(path), &pr.token);
|
||||
detail::vdpf::fold_output_binding(pr.token, key);
|
||||
return out;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// eval_interval(target, key, from, to [, buf [, memo]])
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <std::size_t I, std::size_t N, typename KeyT, typename LaneT,
|
||||
typename OutputBuffer, typename IntervalMemoizer>
|
||||
typename OutputBuffer, typename IntervalMemoizer,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_interval(out_t<I, N>, const KeyT & key, LaneT from, LaneT to,
|
||||
OutputBuffer && outbuf, IntervalMemoizer && memo)
|
||||
{
|
||||
|
|
@ -130,8 +187,10 @@ auto eval_interval(out_t<I, N>, const KeyT & key, LaneT from, LaneT to,
|
|||
}
|
||||
}
|
||||
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <std::size_t I, std::size_t N, typename KeyT, typename LaneT,
|
||||
typename OutputBuffer>
|
||||
typename OutputBuffer,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_interval(out_t<I, N>, const KeyT & key, LaneT from, LaneT to,
|
||||
OutputBuffer && outbuf)
|
||||
{
|
||||
|
|
@ -148,7 +207,9 @@ auto eval_interval(out_t<I, N>, const KeyT & key, LaneT from, LaneT to,
|
|||
}
|
||||
}
|
||||
|
||||
template <std::size_t I, std::size_t N, typename KeyT, typename LaneT>
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <std::size_t I, std::size_t N, typename KeyT, typename LaneT,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_interval(out_t<I, N>, const KeyT & key, LaneT from, LaneT to)
|
||||
{
|
||||
if constexpr (is_multilevel_key_v<KeyT>)
|
||||
|
|
@ -162,8 +223,10 @@ auto eval_interval(out_t<I, N>, const KeyT & key, LaneT from, LaneT to)
|
|||
}
|
||||
}
|
||||
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <typename Beta = uint64_t, typename KeyT, typename LaneT,
|
||||
typename OutputBuffer>
|
||||
typename OutputBuffer,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
void eval_interval(cmp_t, const KeyT & key, LaneT from, LaneT to,
|
||||
OutputBuffer && outbuf)
|
||||
{
|
||||
|
|
@ -171,8 +234,25 @@ void eval_interval(cmp_t, const KeyT & key, LaneT from, LaneT to,
|
|||
std::forward<OutputBuffer>(outbuf));
|
||||
}
|
||||
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <typename Beta = uint64_t, typename KeyT, typename LaneT,
|
||||
typename OutputBuffer, typename IntervalMemoizer>
|
||||
typename OutputBuffer,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
void eval_interval(cmp_t, const KeyT & key, LaneT from, LaneT to,
|
||||
OutputBuffer && outbuf, prove_ref pr)
|
||||
{
|
||||
static_assert(KeyT::is_verifiable,
|
||||
"eval_interval(cmp, ..., prove(π)): key must carry dpf::verifiable");
|
||||
detail::vdpf::init_proof(pr.token, key);
|
||||
detail::incr::eval_cmp_interval_impl<Beta>(key, from, to,
|
||||
std::forward<OutputBuffer>(outbuf), &pr.token);
|
||||
detail::vdpf::fold_output_binding(pr.token, key);
|
||||
}
|
||||
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <typename Beta = uint64_t, typename KeyT, typename LaneT,
|
||||
typename OutputBuffer, typename IntervalMemoizer,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
void eval_interval(cmp_t, const KeyT & key, LaneT from, LaneT to,
|
||||
OutputBuffer && outbuf, IntervalMemoizer && memo)
|
||||
{
|
||||
|
|
@ -181,18 +261,51 @@ void eval_interval(cmp_t, const KeyT & key, LaneT from, LaneT to,
|
|||
std::forward<IntervalMemoizer>(memo));
|
||||
}
|
||||
|
||||
template <typename Beta = uint64_t, typename KeyT, typename LaneT>
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <typename Beta = uint64_t, typename KeyT, typename LaneT,
|
||||
typename OutputBuffer, typename IntervalMemoizer,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
void eval_interval(cmp_t, const KeyT & key, LaneT from, LaneT to,
|
||||
OutputBuffer && outbuf, IntervalMemoizer && memo, prove_ref pr)
|
||||
{
|
||||
static_assert(KeyT::is_verifiable,
|
||||
"eval_interval(cmp, ..., prove(π)): key must carry dpf::verifiable");
|
||||
detail::vdpf::init_proof(pr.token, key);
|
||||
detail::incr::eval_cmp_interval_impl<Beta>(key, from, to,
|
||||
std::forward<OutputBuffer>(outbuf),
|
||||
std::forward<IntervalMemoizer>(memo), &pr.token);
|
||||
detail::vdpf::fold_output_binding(pr.token, key);
|
||||
}
|
||||
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <typename Beta = uint64_t, typename KeyT, typename LaneT,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_interval(cmp_t, const KeyT & key, LaneT from, LaneT to)
|
||||
{
|
||||
return detail::incr::eval_cmp_interval_impl<Beta>(key, from, to);
|
||||
}
|
||||
|
||||
/// \complexity O(L) interior traversals and O(L) workspace in the basic memoizer. L is the number of leaf nodes covering the closed interval (`get_nodes_at_level` at `depth`). Level k expands `(to >> (n-k)) - (from >> (n-k)) + 1` nodes; those counts sum to Θ(L). The output buffer holds one slot per input in the interval. n is `depth`.
|
||||
template <typename Beta = uint64_t, typename KeyT, typename LaneT,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_interval(cmp_t, const KeyT & key, LaneT from, LaneT to, prove_ref pr)
|
||||
{
|
||||
static_assert(KeyT::is_verifiable,
|
||||
"eval_interval(cmp, ..., prove(π)): key must carry dpf::verifiable");
|
||||
detail::vdpf::init_proof(pr.token, key);
|
||||
auto out = detail::incr::eval_cmp_interval_impl<Beta>(key, from, to, &pr.token);
|
||||
detail::vdpf::fold_output_binding(pr.token, key);
|
||||
return out;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// eval_full(target, key [, …])
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// \complexity Same expansion as `eval_interval` on the whole domain. L = 2^{n - lg(outputs_per_leaf)} leaf nodes, n = `depth`. Time Θ(L) interior traversals. The output buffer stores one slot per domain point (2^n).
|
||||
template <std::size_t I, std::size_t N, typename KeyT,
|
||||
typename OutputBuffer, typename IntervalMemoizer>
|
||||
typename OutputBuffer, typename IntervalMemoizer,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_full(out_t<I, N>, const KeyT & key, OutputBuffer && outbuf,
|
||||
IntervalMemoizer && memo)
|
||||
{
|
||||
|
|
@ -210,7 +323,9 @@ auto eval_full(out_t<I, N>, const KeyT & key, OutputBuffer && outbuf,
|
|||
}
|
||||
}
|
||||
|
||||
template <std::size_t I, std::size_t N, typename KeyT>
|
||||
/// \complexity Same expansion as `eval_interval` on the whole domain. L = 2^{n - lg(outputs_per_leaf)} leaf nodes, n = `depth`. Time Θ(L) interior traversals. The output buffer stores one slot per domain point (2^n).
|
||||
template <std::size_t I, std::size_t N, typename KeyT,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_full(out_t<I, N>, const KeyT & key)
|
||||
{
|
||||
if constexpr (is_multilevel_key_v<KeyT>)
|
||||
|
|
@ -224,7 +339,9 @@ auto eval_full(out_t<I, N>, const KeyT & key)
|
|||
}
|
||||
}
|
||||
|
||||
template <typename Beta = uint64_t, typename KeyT>
|
||||
/// \complexity Same expansion as `eval_interval` on the whole domain. L = 2^{n - lg(outputs_per_leaf)} leaf nodes, n = `depth`. Time Θ(L) interior traversals. The output buffer stores one slot per domain point (2^n).
|
||||
template <typename Beta = uint64_t, typename KeyT,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_full(cmp_t, const KeyT & key)
|
||||
{
|
||||
if (!key.has_cmp())
|
||||
|
|
@ -238,13 +355,36 @@ auto eval_full(cmp_t, const KeyT & key)
|
|||
return detail::incr::eval_cmp_interval_impl<Beta>(key, lo, hi);
|
||||
}
|
||||
|
||||
/// \complexity Same expansion as `eval_interval` on the whole domain. L = 2^{n - lg(outputs_per_leaf)} leaf nodes, n = `depth`. Time Θ(L) interior traversals. The output buffer stores one slot per domain point (2^n).
|
||||
template <typename Beta = uint64_t, typename KeyT,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_full(cmp_t, const KeyT & key, prove_ref pr)
|
||||
{
|
||||
static_assert(KeyT::is_verifiable,
|
||||
"eval_full(cmp, ..., prove(π)): key must carry dpf::verifiable");
|
||||
if (!key.has_cmp())
|
||||
throw std::invalid_argument("eval_full(cmp): no comparison channel");
|
||||
using lane_t = typename KeyT::integral_type;
|
||||
const auto nbits = static_cast<std::size_t>(key.cmp().nbits);
|
||||
const lane_t lo = 0;
|
||||
const lane_t hi = (nbits >= 8 * sizeof(lane_t))
|
||||
? static_cast<lane_t>(~lane_t{0})
|
||||
: static_cast<lane_t>((lane_t{1} << nbits) - 1);
|
||||
detail::vdpf::init_proof(pr.token, key);
|
||||
auto out = detail::incr::eval_cmp_interval_impl<Beta>(key, lo, hi, &pr.token);
|
||||
detail::vdpf::fold_output_binding(pr.token, key);
|
||||
return out;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// eval_sequence(target, key, begin, end, buf [, path])
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// \complexity O(n k) interior traversals in the worst case and O(k) node workspace. k is the number of listed points and n is `depth`. The breadth-first buffer is 2k nodes, so each level traverses at most one node per point. Shared prefixes do fewer traversals. A recipe memoizer instead stores O(recipe leaf nodes) (see that memoizer).
|
||||
template <std::size_t I, std::size_t N, typename KeyT, typename ForwardIterator,
|
||||
typename OutputBuffer,
|
||||
typename PathMemoizer = basic_path_memoizer<KeyT>>
|
||||
typename PathMemoizer = basic_path_memoizer<KeyT>,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_sequence(out_t<I, N>, const KeyT & key, ForwardIterator begin,
|
||||
ForwardIterator end, OutputBuffer && outbuf,
|
||||
PathMemoizer && path = PathMemoizer{})
|
||||
|
|
@ -263,9 +403,11 @@ auto eval_sequence(out_t<I, N>, const KeyT & key, ForwardIterator begin,
|
|||
}
|
||||
}
|
||||
|
||||
/// \complexity O(n k) interior traversals in the worst case and O(k) node workspace. k is the number of listed points and n is `depth`. The breadth-first buffer is 2k nodes, so each level traverses at most one node per point. Shared prefixes do fewer traversals. A recipe memoizer instead stores O(recipe leaf nodes) (see that memoizer).
|
||||
template <typename Beta = uint64_t, typename KeyT, typename ForwardIterator,
|
||||
typename OutputBuffer,
|
||||
typename PathMemoizer = basic_path_memoizer<KeyT>>
|
||||
typename PathMemoizer = basic_path_memoizer<KeyT>,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
void eval_sequence(cmp_t, const KeyT & key, ForwardIterator begin,
|
||||
ForwardIterator end, OutputBuffer && outbuf,
|
||||
PathMemoizer && path = PathMemoizer{})
|
||||
|
|
@ -275,24 +417,45 @@ void eval_sequence(cmp_t, const KeyT & key, ForwardIterator begin,
|
|||
std::forward<PathMemoizer>(path));
|
||||
}
|
||||
|
||||
/// \complexity O(n k) interior traversals in the worst case and O(k) node workspace. k is the number of listed points and n is `depth`. The breadth-first buffer is 2k nodes, so each level traverses at most one node per point. Shared prefixes do fewer traversals. A recipe memoizer instead stores O(recipe leaf nodes) (see that memoizer).
|
||||
template <typename Beta = uint64_t, typename KeyT, typename ForwardIterator,
|
||||
typename OutputBuffer,
|
||||
typename PathMemoizer = basic_path_memoizer<KeyT>,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
void eval_sequence(cmp_t, const KeyT & key, ForwardIterator begin,
|
||||
ForwardIterator end, OutputBuffer && outbuf, prove_ref pr,
|
||||
PathMemoizer && path = PathMemoizer{})
|
||||
{
|
||||
static_assert(KeyT::is_verifiable,
|
||||
"eval_sequence(cmp, ..., prove(π)): key must carry dpf::verifiable");
|
||||
detail::vdpf::init_proof(pr.token, key);
|
||||
detail::incr::eval_cmp_sequence_impl<Beta>(key, begin, end,
|
||||
std::forward<OutputBuffer>(outbuf),
|
||||
std::forward<PathMemoizer>(path), &pr.token);
|
||||
detail::vdpf::fold_output_binding(pr.token, key);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// make_output_buffer(target, …)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
template <typename Beta = uint64_t, typename KeyT>
|
||||
template <typename Beta = uint64_t, typename KeyT,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto make_output_buffer(cmp_t, const KeyT & key, std::size_t n)
|
||||
{
|
||||
return detail::incr::make_output_buffer_for_cmp_impl<Beta>(key, n);
|
||||
}
|
||||
|
||||
template <typename Beta = uint64_t, typename KeyT, typename LaneT>
|
||||
template <typename Beta = uint64_t, typename KeyT, typename LaneT,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto make_output_buffer(cmp_t, const KeyT & key, LaneT from, LaneT to)
|
||||
{
|
||||
return detail::incr::make_output_buffer_for_cmp_interval_impl<Beta>(
|
||||
key, from, to);
|
||||
}
|
||||
|
||||
template <std::size_t I, std::size_t N, typename KeyT, typename LaneT>
|
||||
template <std::size_t I, std::size_t N, typename KeyT, typename LaneT,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto make_output_buffer(out_t<I, N>, const KeyT & key, LaneT from, LaneT to)
|
||||
{
|
||||
constexpr auto pref = detail::resolved_out_prefix<I, N, KeyT>();
|
||||
|
|
@ -426,6 +589,7 @@ auto eval_out_inner_product_impl(const KeyT & dpf, LaneT from, LaneT to,
|
|||
for (std::size_t j = 0; j < count; ++j)
|
||||
{
|
||||
auto leaf = dpf.template traverse_exterior<I>(nodes[j]);
|
||||
detail::incr::absorb_public_addend_all_lanes<I>(dpf, leaf);
|
||||
acc.mac(leaf, (start + j) * opl, opl, weights);
|
||||
}
|
||||
start += seg.count;
|
||||
|
|
@ -436,7 +600,7 @@ auto eval_out_inner_product_impl(const KeyT & dpf, LaneT from, LaneT to,
|
|||
template <typename Beta = uint64_t, typename KeyT, typename LaneT,
|
||||
typename Weights>
|
||||
Beta eval_cmp_inner_product_impl(const KeyT & dpf, LaneT from, LaneT to,
|
||||
Weights && weights)
|
||||
Weights && weights, proof_token * pi = nullptr)
|
||||
{
|
||||
if (!dpf.has_cmp())
|
||||
throw std::invalid_argument("cmp inner product: no comparison channel");
|
||||
|
|
@ -459,7 +623,7 @@ Beta eval_cmp_inner_product_impl(const KeyT & dpf, LaneT from, LaneT to,
|
|||
const std::size_t levels = unwrap_party_key_t<KeyT>::cmp_block > 0
|
||||
? unwrap_party_key_t<KeyT>::cmp_h : nbits;
|
||||
detail::incr::eval_cmp_interval_impl_interior(dpf, a, cmp_exclusive_end(b),
|
||||
nbits, memo, levels);
|
||||
nbits, memo, levels, pi);
|
||||
|
||||
uint64_t dot = 0;
|
||||
for (std::size_t i = 0; i < count; ++i)
|
||||
|
|
@ -486,9 +650,11 @@ Beta eval_cmp_inner_product_impl(const KeyT & dpf, LaneT from, LaneT to,
|
|||
} // namespace incr
|
||||
} // namespace detail
|
||||
|
||||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||||
template <std::size_t I, std::size_t N, typename KeyT, typename LaneT,
|
||||
typename Weights, typename IntervalMemoizer,
|
||||
std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true>
|
||||
std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_inner_product(out_t<I, N>, const KeyT & key, LaneT from, LaneT to,
|
||||
Weights && weights, IntervalMemoizer && memo)
|
||||
{
|
||||
|
|
@ -498,20 +664,24 @@ auto eval_inner_product(out_t<I, N>, const KeyT & key, LaneT from, LaneT to,
|
|||
std::forward<IntervalMemoizer>(memo));
|
||||
}
|
||||
|
||||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||||
template <std::size_t I, std::size_t N, typename KeyT, typename LaneT,
|
||||
typename Weights,
|
||||
std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true>
|
||||
std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_inner_product(out_t<I, N>, const KeyT & key, LaneT from, LaneT to,
|
||||
Weights && weights)
|
||||
{
|
||||
auto memo = make_basic_interval_memoizer<KeyT, I>(from, to);
|
||||
auto memo = make_basic_interval_memoizer<KeyT, I>(key, from, to);
|
||||
constexpr auto pref = detail::resolved_out_prefix<I, N, KeyT>();
|
||||
return detail::incr::eval_out_inner_product_impl<pref, I>(key, from, to,
|
||||
std::forward<Weights>(weights), memo);
|
||||
}
|
||||
|
||||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||||
template <typename Beta = uint64_t, typename KeyT, typename LaneT,
|
||||
typename Weights>
|
||||
typename Weights,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
Beta eval_inner_product(cmp_t, const KeyT & key, LaneT from, LaneT to,
|
||||
Weights && weights)
|
||||
{
|
||||
|
|
@ -519,6 +689,106 @@ Beta eval_inner_product(cmp_t, const KeyT & key, LaneT from, LaneT to,
|
|||
std::forward<Weights>(weights));
|
||||
}
|
||||
|
||||
/// @brief Comparison inner product over the whole comparison domain.
|
||||
/// @details `sum_x [x satisfies cmp] * weights[x]` (as complementary halves),
|
||||
/// the full-domain form of `eval_inner_product(cmp, key, lo, hi, w)`.
|
||||
/// Pair the two parties' results with `reconstruct_cmp_halves`.
|
||||
/// Waldo's private-threshold aggregate is this one call.
|
||||
/// \complexity Same expansion as `eval_full` on the comparison domain, plus a
|
||||
/// multiply-add per point into an `O(1)` accumulator.
|
||||
template <typename Beta = uint64_t, typename KeyT, typename Weights,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
Beta eval_full_inner_product(cmp_t, const KeyT & key, Weights && weights)
|
||||
{
|
||||
if (!key.has_cmp())
|
||||
throw std::invalid_argument(
|
||||
"eval_full_inner_product(cmp): no comparison channel");
|
||||
using lane_t = typename KeyT::integral_type;
|
||||
const auto nbits = static_cast<std::size_t>(key.cmp().nbits);
|
||||
const lane_t lo = 0;
|
||||
const lane_t hi = (nbits >= 8 * sizeof(lane_t))
|
||||
? static_cast<lane_t>(~lane_t{0})
|
||||
: static_cast<lane_t>((lane_t{1} << nbits) - 1);
|
||||
return eval_inner_product<Beta>(cmp, key, lo, hi,
|
||||
std::forward<Weights>(weights));
|
||||
}
|
||||
|
||||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||||
template <typename Beta = uint64_t, typename KeyT, typename LaneT,
|
||||
typename Weights,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
Beta eval_inner_product(cmp_t, const KeyT & key, LaneT from, LaneT to,
|
||||
Weights && weights, prove_ref pr)
|
||||
{
|
||||
static_assert(KeyT::is_verifiable,
|
||||
"eval_inner_product(cmp, ..., prove(π)): key must carry dpf::verifiable");
|
||||
detail::vdpf::init_proof(pr.token, key);
|
||||
auto out = detail::incr::eval_cmp_inner_product_impl<Beta>(key, from, to,
|
||||
std::forward<Weights>(weights), &pr.token);
|
||||
detail::vdpf::fold_output_binding(pr.token, key);
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief Initialise `pr.token` and fold `[from, to]` once per cmp BFS node.
|
||||
template <typename KeyT, typename LaneT,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
void prove_cmp_interval(const KeyT & key, LaneT from, LaneT to, prove_ref pr)
|
||||
{
|
||||
static_assert(KeyT::is_verifiable,
|
||||
"prove_cmp_interval: key must carry dpf::verifiable");
|
||||
detail::vdpf::init_proof(pr.token, key);
|
||||
detail::incr::prove_fold_cmp_interval(key, from, to, pr.token);
|
||||
detail::vdpf::fold_output_binding(pr.token, key);
|
||||
}
|
||||
|
||||
/// @brief Initialise `pr.token` and fold the full comparison domain.
|
||||
template <typename KeyT,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
void prove_cmp_full(const KeyT & key, prove_ref pr)
|
||||
{
|
||||
static_assert(KeyT::is_verifiable,
|
||||
"prove_cmp_full: key must carry dpf::verifiable");
|
||||
if (!key.has_cmp())
|
||||
throw std::invalid_argument("prove_cmp_full: no comparison channel");
|
||||
using lane_t = typename KeyT::integral_type;
|
||||
const auto nbits = static_cast<std::size_t>(key.cmp().nbits);
|
||||
const lane_t lo = 0;
|
||||
const lane_t hi = (nbits >= 8 * sizeof(lane_t))
|
||||
? static_cast<lane_t>(~lane_t{0})
|
||||
: static_cast<lane_t>((lane_t{1} << nbits) - 1);
|
||||
prove_cmp_interval(key, lo, hi, pr);
|
||||
}
|
||||
|
||||
/// @brief Fold a sorted sequence into `pr` via cmp interval-run covers.
|
||||
template <typename KeyT, typename ForwardIterator,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
void prove_cmp_sequence(const KeyT & key, ForwardIterator begin,
|
||||
ForwardIterator end, prove_ref pr)
|
||||
{
|
||||
static_assert(KeyT::is_verifiable,
|
||||
"prove_cmp_sequence: key must carry dpf::verifiable");
|
||||
if (HEDLEY_UNLIKELY(begin != end && !std::is_sorted(begin, end)))
|
||||
throw std::runtime_error("list must be sorted");
|
||||
detail::vdpf::init_proof(pr.token, key);
|
||||
using lane_t = typename KeyT::integral_type;
|
||||
for (auto it = begin; it != end; )
|
||||
{
|
||||
const auto run_from = static_cast<lane_t>(*it);
|
||||
auto run_to = run_from;
|
||||
++it;
|
||||
while (it != end)
|
||||
{
|
||||
const auto next = static_cast<lane_t>(*it);
|
||||
if (next != static_cast<lane_t>(run_to + lane_t{1}))
|
||||
break;
|
||||
run_to = next;
|
||||
++it;
|
||||
}
|
||||
detail::incr::prove_fold_cmp_interval(key, run_from, run_to, pr.token);
|
||||
}
|
||||
detail::vdpf::fold_output_binding(pr.token, key);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// eval_sequence_breadth_first(out<I>, key, begin, end [, outbuf])
|
||||
//
|
||||
|
|
@ -621,6 +891,8 @@ void eval_out_sequence_breadth_first_impl(const KeyT & dpf,
|
|||
auto leaf = dpf.template traverse_exterior<I>(buf[j]);
|
||||
const std::size_t off =
|
||||
static_cast<std::size_t>(static_cast<input_type>(*curr) & (opl - 1));
|
||||
detail::incr::absorb_public_addend_lane<I>(dpf, leaf,
|
||||
static_cast<input_type>(off));
|
||||
auto v = dpf::extract_leaf<exterior_node, output_type>(leaf, off);
|
||||
if constexpr (is_party_key_v<KeyT>)
|
||||
outbuf[i] = subtractive_share<output_type, party_of_v<KeyT>>::from_raw(v);
|
||||
|
|
@ -633,9 +905,11 @@ void eval_out_sequence_breadth_first_impl(const KeyT & dpf,
|
|||
} // namespace incr
|
||||
} // namespace detail
|
||||
|
||||
/// \complexity O(n k) interior traversals in the worst case and O(k) node workspace. k is the number of listed points and n is `depth`. The breadth-first buffer is 2k nodes, so each level traverses at most one node per point. Shared prefixes do fewer traversals. A recipe memoizer instead stores O(recipe leaf nodes) (see that memoizer).
|
||||
template <std::size_t I, std::size_t N, typename KeyT,
|
||||
typename ForwardIterator, typename OutputBuffer,
|
||||
std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true>
|
||||
std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
void eval_sequence_breadth_first(out_t<I, N>, const KeyT & key,
|
||||
ForwardIterator begin, ForwardIterator end, OutputBuffer && outbuf)
|
||||
{
|
||||
|
|
@ -644,9 +918,11 @@ void eval_sequence_breadth_first(out_t<I, N>, const KeyT & key,
|
|||
std::forward<OutputBuffer>(outbuf));
|
||||
}
|
||||
|
||||
/// \complexity O(n k) interior traversals in the worst case and O(k) node workspace. k is the number of listed points and n is `depth`. The breadth-first buffer is 2k nodes, so each level traverses at most one node per point. Shared prefixes do fewer traversals. A recipe memoizer instead stores O(recipe leaf nodes) (see that memoizer).
|
||||
template <std::size_t I, std::size_t N, typename KeyT,
|
||||
typename ForwardIterator,
|
||||
std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true>
|
||||
std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto eval_sequence_breadth_first(out_t<I, N>, const KeyT & key,
|
||||
ForwardIterator begin, ForwardIterator end)
|
||||
{
|
||||
|
|
@ -668,7 +944,8 @@ auto eval_sequence_breadth_first(out_t<I, N>, const KeyT & key,
|
|||
/// @param end the iterator past the last query
|
||||
/// @return the constructed object
|
||||
template <std::size_t I, std::size_t N, typename KeyT, typename ForwardIterator,
|
||||
std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true>
|
||||
std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<KeyT>>, int> = 0>
|
||||
auto make_sequence_recipe(out_t<I, N>, const KeyT & key, ForwardIterator begin,
|
||||
ForwardIterator end)
|
||||
{
|
||||
|
|
|
|||
323
include/dpf/eval_until.hpp
Normal file
323
include/dpf/eval_until.hpp
Normal file
|
|
@ -0,0 +1,323 @@
|
|||
/// @file dpf/eval_until.hpp
|
||||
/// @brief Prefix-resuming evaluation for multilevel / incremental DPF keys.
|
||||
/// @details Google's incremental-DPF `EvaluateUntil(level, prefixes, ctx)`
|
||||
/// (Poplar / private heavy hitters, ePrint 2021/017) walks only the
|
||||
/// live prefixes. `eval_prefixes` always restarts at the root and
|
||||
/// materializes all `2^N` nodes; a path memoizer resumes one path.
|
||||
/// `idpf_eval_ctx` keeps the interior node under each live prefix.
|
||||
/// `eval_until(ctx, level, prefixes)` returns the output shares at
|
||||
/// that hierarchy level for those prefixes only, then updates the
|
||||
/// context. Empty `prefixes` on a fresh context (level 0) leaves the
|
||||
/// root seed in place. Each later prefix must extend a prefix from
|
||||
/// the previous call.
|
||||
/// @see dpf/eval_walk.hpp (`eval_prefixes`), dpf/idpf_agg.hpp
|
||||
/// @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_DPF_EVAL_UNTIL_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_EVAL_UNTIL_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <stdexcept>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/dpf_key.hpp"
|
||||
#include "dpf/eval_common.hpp"
|
||||
#include "dpf/incremental.hpp"
|
||||
#include "dpf/placement.hpp"
|
||||
#include "dpf/secret_share.hpp"
|
||||
#include "dpf/utils.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
/// @brief Live-prefix evaluation context for a multilevel / idpf key.
|
||||
/// @details Hierarchy level 0 holds only the root seed. After
|
||||
/// `eval_until(..., L, prefixes)`, `level()` is `L` and
|
||||
/// `node_count()` equals `prefixes.size()`.
|
||||
/// \complexity O(|prefixes|) stored nodes (not O(2^level)).
|
||||
template <typename KeyT>
|
||||
class idpf_eval_ctx
|
||||
{
|
||||
public:
|
||||
using key_type = unwrap_party_key_t<KeyT>;
|
||||
using input_type = typename key_type::input_type;
|
||||
using node_type = typename key_type::interior_node;
|
||||
|
||||
explicit idpf_eval_ctx(const KeyT & key)
|
||||
: key_{&key}, level_{0}, prefixes_{}, nodes_{}
|
||||
{
|
||||
nodes_.push_back(static_cast<const key_type &>(key).root());
|
||||
// Compact parent of every length-1 prefix is the empty prefix `0`.
|
||||
prefixes_.push_back(input_type{0});
|
||||
}
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
const KeyT & key() const noexcept { return *key_; }
|
||||
|
||||
/// @brief Last hierarchy level that `eval_until` wrote (0 = root only).
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
std::size_t level() const noexcept { return level_; }
|
||||
|
||||
/// @brief Number of saved interior nodes (one per live prefix).
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
std::size_t node_count() const noexcept { return nodes_.size(); }
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
const std::vector<input_type> & prefixes() const noexcept
|
||||
{
|
||||
return prefixes_;
|
||||
}
|
||||
|
||||
/// @brief Drop every saved prefix except `prefix` (must be live).
|
||||
void retain(input_type prefix)
|
||||
{
|
||||
for (std::size_t i = 0; i < prefixes_.size(); ++i)
|
||||
{
|
||||
if (prefixes_[i] == prefix)
|
||||
{
|
||||
prefixes_ = {prefix};
|
||||
nodes_ = {nodes_[i]};
|
||||
return;
|
||||
}
|
||||
}
|
||||
throw std::invalid_argument(
|
||||
"idpf_eval_ctx::retain: prefix is not live in this context");
|
||||
}
|
||||
|
||||
private:
|
||||
template <typename K, typename PrefRange>
|
||||
friend auto eval_until(idpf_eval_ctx<K> & ctx, std::size_t level,
|
||||
PrefRange && prefixes);
|
||||
|
||||
const KeyT * key_;
|
||||
std::size_t level_;
|
||||
std::vector<input_type> prefixes_;
|
||||
std::vector<node_type> nodes_;
|
||||
};
|
||||
|
||||
namespace detail
|
||||
{
|
||||
namespace eval_until_detail
|
||||
{
|
||||
|
||||
template <typename KeyT>
|
||||
HEDLEY_NO_THROW
|
||||
constexpr std::size_t slot_for_prefix(std::size_t prefix_len) noexcept
|
||||
{
|
||||
for (std::size_t i = 0; i < KeyT::num_outputs; ++i)
|
||||
{
|
||||
if (KeyT::meta[i].prefix == prefix_len)
|
||||
return i;
|
||||
}
|
||||
return static_cast<std::size_t>(-1);
|
||||
}
|
||||
|
||||
template <typename KeyT>
|
||||
HEDLEY_NO_THROW
|
||||
constexpr std::size_t tree_level_for_prefix(std::size_t prefix_len) noexcept
|
||||
{
|
||||
const auto i = slot_for_prefix<KeyT>(prefix_len);
|
||||
return (i == static_cast<std::size_t>(-1)) ? static_cast<std::size_t>(-1)
|
||||
: KeyT::meta[i].tree_level;
|
||||
}
|
||||
|
||||
/// @brief Walk `from_tree_level` → `to_tree_level` along the high bits of `x`.
|
||||
/// \complexity O(to − from) interior traversals.
|
||||
template <typename KeyT, typename Node>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
Node walk_interior(const KeyT & key, Node node, typename KeyT::input_type x,
|
||||
std::size_t from_tree_level, std::size_t to_tree_level)
|
||||
{
|
||||
using key_type = KeyT;
|
||||
if (to_tree_level <= from_tree_level)
|
||||
return node;
|
||||
auto level_index = from_tree_level + 1;
|
||||
auto mask = key.msb_mask >> (level_index - 1);
|
||||
for (; level_index <= to_tree_level; ++level_index, mask >>= 1)
|
||||
{
|
||||
const bool bit = !!(mask & x);
|
||||
auto cw = key.correction_word(level_index - 1, bit);
|
||||
const bool is_last = key_type::tree::is_last_level(level_index - 1,
|
||||
key.depth);
|
||||
node = key_type::traverse_interior(node, cw, bit, is_last);
|
||||
}
|
||||
return node;
|
||||
}
|
||||
|
||||
template <std::size_t I, typename KeyT, typename Node, typename InputT>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto exterior_at(const KeyT & key, const Node & node, InputT domain_x)
|
||||
{
|
||||
using key_type = unwrap_party_key_t<KeyT>;
|
||||
using output_type = typename key_type::template concrete_output_type<I>;
|
||||
constexpr auto N = key_type::meta[I].prefix;
|
||||
// `traverse_exterior` lives on the underlying key; party_key inherits it.
|
||||
auto leaf = key.template traverse_exterior<I>(node);
|
||||
auto lane_x = detail::incr::lane_input(domain_x, N, key_type::input_bits);
|
||||
detail::incr::absorb_public_addend_lane<I>(
|
||||
static_cast<const key_type &>(key), leaf, lane_x);
|
||||
// Pass the party-tagged key type so the share carries the party id.
|
||||
return *make_eval_dpf_output<KeyT, output_type>(leaf, lane_x);
|
||||
}
|
||||
|
||||
template <typename KeyT, typename Node, typename InputT, typename Out,
|
||||
std::size_t... Is>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
bool exterior_dispatch(std::size_t slot, const KeyT & key, const Node & node,
|
||||
InputT domain_x, Out & out, std::index_sequence<Is...>)
|
||||
{
|
||||
return ((Is == slot
|
||||
? (out = exterior_at<Is>(key, node, domain_x), true)
|
||||
: false)
|
||||
|| ...);
|
||||
}
|
||||
|
||||
template <typename KeyT>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto share_type_tag()
|
||||
{
|
||||
using key_type = unwrap_party_key_t<KeyT>;
|
||||
using output_type = typename key_type::template concrete_output_type<0>;
|
||||
return eval_leaf_result_t<KeyT, output_type>{};
|
||||
}
|
||||
|
||||
} // namespace eval_until_detail
|
||||
} // namespace detail
|
||||
|
||||
/// @brief Evaluate output shares at hierarchy `level` for `prefixes` only.
|
||||
/// @details Each prefix is a compact integer in `[0, 2^level)`. The first call
|
||||
/// may pass an empty range at level 0 to keep the root seed. Later
|
||||
/// calls require `level > ctx.level()` and every prefix to extend one
|
||||
/// saved parent. Returns one share per prefix, in the same order.
|
||||
/// Opened values match `eval_point(out<I,N>, key, domain)` for the
|
||||
/// same prefix length (see `eval_until_test`).
|
||||
/// \complexity O(|prefixes| · (level − ctx.level())) interior traversals.
|
||||
/// Saved nodes stay O(|prefixes|), not O(2^level).
|
||||
/// @see ePrint 2021/017 (Poplar EvaluateUntil), ePrint 2024/1190 (I-DPF agg)
|
||||
template <typename KeyT, typename PrefRange>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_until(idpf_eval_ctx<KeyT> & ctx, std::size_t level,
|
||||
PrefRange && prefixes)
|
||||
{
|
||||
using key_type = typename idpf_eval_ctx<KeyT>::key_type;
|
||||
using input_type = typename idpf_eval_ctx<KeyT>::input_type;
|
||||
using node_type = typename idpf_eval_ctx<KeyT>::node_type;
|
||||
using share_type = decltype(detail::eval_until_detail::share_type_tag<KeyT>());
|
||||
static_assert(is_multilevel_key_v<key_type>,
|
||||
"eval_until requires a multilevel / incremental DPF key");
|
||||
|
||||
const KeyT & key_ref = ctx.key();
|
||||
const key_type & key = static_cast<const key_type &>(key_ref);
|
||||
|
||||
std::vector<input_type> pref_list;
|
||||
for (auto && p : prefixes)
|
||||
pref_list.push_back(static_cast<input_type>(p));
|
||||
|
||||
if (pref_list.empty())
|
||||
{
|
||||
if (ctx.level_ != 0 || level != 0)
|
||||
{
|
||||
throw std::invalid_argument(
|
||||
"eval_until: empty prefixes only valid on a fresh context at level 0");
|
||||
}
|
||||
return std::vector<share_type>{};
|
||||
}
|
||||
|
||||
if (level == 0)
|
||||
{
|
||||
throw std::invalid_argument(
|
||||
"eval_until: level 0 has only the root; pass a positive hierarchy level");
|
||||
}
|
||||
if (level <= ctx.level_)
|
||||
{
|
||||
throw std::invalid_argument(
|
||||
"eval_until: level must strictly advance past the context level");
|
||||
}
|
||||
if (level > key_type::input_bits)
|
||||
{
|
||||
throw std::invalid_argument(
|
||||
"eval_until: level exceeds the input bit length");
|
||||
}
|
||||
|
||||
const auto slot = detail::eval_until_detail::slot_for_prefix<key_type>(level);
|
||||
if (slot == static_cast<std::size_t>(-1))
|
||||
{
|
||||
throw std::invalid_argument(
|
||||
"eval_until: key has no output planted at that prefix length");
|
||||
}
|
||||
const auto to_tree =
|
||||
detail::eval_until_detail::tree_level_for_prefix<key_type>(level);
|
||||
const std::size_t from_tree = (ctx.level_ == 0)
|
||||
? 0
|
||||
: detail::eval_until_detail::tree_level_for_prefix<key_type>(ctx.level_);
|
||||
|
||||
constexpr auto bits = key_type::input_bits;
|
||||
const auto parent_shift = level - ctx.level_;
|
||||
|
||||
std::vector<node_type> new_nodes;
|
||||
std::vector<share_type> shares;
|
||||
new_nodes.reserve(pref_list.size());
|
||||
shares.reserve(pref_list.size());
|
||||
|
||||
for (auto p : pref_list)
|
||||
{
|
||||
if (level < bits && (static_cast<std::uint64_t>(p) >> level) != 0)
|
||||
{
|
||||
throw std::invalid_argument(
|
||||
"eval_until: prefix does not fit in the requested bit length");
|
||||
}
|
||||
const input_type parent = static_cast<input_type>(p >> parent_shift);
|
||||
std::size_t parent_idx = static_cast<std::size_t>(-1);
|
||||
for (std::size_t i = 0; i < ctx.prefixes_.size(); ++i)
|
||||
{
|
||||
if (ctx.prefixes_[i] == parent)
|
||||
{
|
||||
parent_idx = i;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if (parent_idx == static_cast<std::size_t>(-1))
|
||||
{
|
||||
throw std::invalid_argument(
|
||||
"eval_until: prefix does not extend a live parent");
|
||||
}
|
||||
|
||||
auto domain_x = static_cast<input_type>(
|
||||
static_cast<std::uint64_t>(p) << (bits - level));
|
||||
domain_x = key.offset_x(domain_x);
|
||||
utils::flip_msb_if_signed_integral(domain_x);
|
||||
|
||||
node_type node = detail::eval_until_detail::walk_interior(key,
|
||||
ctx.nodes_[parent_idx], domain_x, from_tree, to_tree);
|
||||
|
||||
share_type share{};
|
||||
const bool ok = detail::eval_until_detail::exterior_dispatch(slot,
|
||||
key_ref, node, domain_x, share,
|
||||
std::make_index_sequence<key_type::num_outputs>{});
|
||||
if (!ok)
|
||||
{
|
||||
throw std::logic_error("eval_until: exterior dispatch missed slot");
|
||||
}
|
||||
new_nodes.push_back(node);
|
||||
shares.push_back(share);
|
||||
}
|
||||
|
||||
ctx.level_ = level;
|
||||
ctx.prefixes_ = std::move(pref_list);
|
||||
ctx.nodes_ = std::move(new_nodes);
|
||||
return shares;
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_EVAL_UNTIL_HPP__
|
||||
449
include/dpf/eval_walk.hpp
Normal file
449
include/dpf/eval_walk.hpp
Normal file
|
|
@ -0,0 +1,449 @@
|
|||
/// @file dpf/eval_walk.hpp
|
||||
/// @brief Walk helpers that fold a small operation into an existing DPF walk.
|
||||
/// @details These are the calls the protocol mockups in
|
||||
/// `examples/applications/` had to build by hand around a walk:
|
||||
/// - `dpf::rotate{s}` weights the walk by `w[(i + s) mod 2^n]`
|
||||
/// (Duoram's read, Pika's lookup) with no second vector.
|
||||
/// - `eval_full_add_into(buf, key)` adds a full-domain expansion
|
||||
/// into a buffer the caller already holds (Prio's histogram,
|
||||
/// Express's mailbox, Duoram's update). A `rotate` overload shifts
|
||||
/// the write; a `sketch` overload folds the audit in the same pass.
|
||||
/// - `dpf::cyclic_shift(buf, s)` rotates a materialized share buffer,
|
||||
/// `new[i] = old[(i - s) mod n]`.
|
||||
/// - `dpf::pack_bit_columns(keys...)` runs the full-domain bit walk
|
||||
/// once per key and packs one integer per row (lane `e` is key `e`),
|
||||
/// the digit BitMore reads when the server count is a power of two.
|
||||
/// - `dpf::mod_bit_columns<ℓ>(keys...)` is that digit modulo `ℓ` when
|
||||
/// the server count is not a power of two. The first key is still
|
||||
/// the low bit. The running residue stays in a byte through `ℓ = 128`
|
||||
/// and in a 16-bit lane through `ℓ = 32768`.
|
||||
/// - `eval_prefixes(out<I,N>, key)` returns the `2^N` prefix shares,
|
||||
/// and `eval_prefix_inner_product(out<I,N>, key, values)` dots them
|
||||
/// with `values` in one walk to depth `N` (Poplar, PRAC).
|
||||
/// - `idpf_eval_ctx` / `eval_until` (see `dpf/eval_until.hpp`) resume
|
||||
/// under a live prefix list instead of materializing `2^N` nodes.
|
||||
/// @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_DPF_EVAL_WALK_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_EVAL_WALK_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <iterator>
|
||||
#include <limits>
|
||||
#include <tuple>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/bitmore_mod.hpp"
|
||||
#include "dpf/eval_target.hpp"
|
||||
#include "dpf/eval_full.hpp"
|
||||
#include "dpf/eval_inner_product.hpp"
|
||||
#include "dpf/eval_unified.hpp"
|
||||
#include "dpf/output_buffer.hpp"
|
||||
#include "dpf/utils.hpp"
|
||||
#include "dpf/verifiable.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
/// @brief Rotation offset applied inside a walk: domain point `i` uses lane
|
||||
/// `(i + shift) mod 2^n`. Pass to `eval_full_inner_product` /
|
||||
/// `eval_full_add_into` so the caller does not build a second vector.
|
||||
struct rotate
|
||||
{
|
||||
std::size_t shift;
|
||||
};
|
||||
|
||||
namespace detail_walk
|
||||
{
|
||||
|
||||
/// @brief Number of input points `2^n` of `KeyT`, as a `std::size_t`.
|
||||
template <typename KeyT>
|
||||
HEDLEY_NO_THROW
|
||||
constexpr std::size_t domain_size() noexcept
|
||||
{
|
||||
constexpr std::size_t bits =
|
||||
utils::bitlength_of_v<typename KeyT::input_type>;
|
||||
static_assert(bits < 8 * sizeof(std::size_t),
|
||||
"walk helper: input domain does not fit a std::size_t index");
|
||||
return std::size_t{1} << bits;
|
||||
}
|
||||
|
||||
/// @brief `w[(i + shift) mod n]`, a read-only rotated view of `w`.
|
||||
template <typename Weights>
|
||||
struct rotated_weights
|
||||
{
|
||||
const Weights & w;
|
||||
std::size_t shift;
|
||||
std::size_t n;
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
decltype(auto) operator[](std::size_t i) const
|
||||
{
|
||||
return w[(i + shift) % n];
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T, typename = void>
|
||||
struct has_raw : std::false_type {};
|
||||
template <typename T>
|
||||
struct has_raw<T, std::void_t<decltype(std::declval<const T &>().raw())>>
|
||||
: std::true_type {};
|
||||
|
||||
/// @brief The group element carried by an eval buffer slot. A one-word share
|
||||
/// (`additive`, `subtractive`, or `additive3`) exposes it through
|
||||
/// `.raw()`. A replicated share has two components and is returned
|
||||
/// as itself. A raw output type is returned unchanged.
|
||||
template <typename T>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto group_value(const T & v)
|
||||
{
|
||||
if constexpr (has_raw<T>::value)
|
||||
return v.raw();
|
||||
else
|
||||
return v;
|
||||
}
|
||||
|
||||
/// @brief Call `fn` on `keys` from the last key down to the first.
|
||||
/// @details `mod_bit_columns` inserts the high bit first. Key 0 stays the low
|
||||
/// bit, matching `pack_bit_columns`.
|
||||
template <typename Fn, typename Tuple, std::size_t... I>
|
||||
void insert_keys_msb_first(Fn && fn, Tuple && keys, std::index_sequence<I...>)
|
||||
{
|
||||
constexpr std::size_t n = sizeof...(I);
|
||||
(fn(std::get<n - 1 - I>(keys)), ...);
|
||||
}
|
||||
|
||||
} // namespace detail_walk
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// dpf::rotate on the paired full-domain inner product
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// @brief `sum_i DPF_I(i) * rows[(i + rot.shift) mod 2^n]` over the whole domain.
|
||||
/// @details The same paired walk as `eval_full_inner_product(paired, key, rows)`,
|
||||
/// but the weight at domain point `i` is read from `rows` rotated by
|
||||
/// `rot.shift`. Duoram's read and Pika's lookup pass the unrotated
|
||||
/// table and this offset instead of materializing a rotated copy.
|
||||
template <std::size_t I = 0, std::size_t... Is, typename DpfKey, typename Rows>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_full_inner_product(paired_t, const DpfKey & dpf, Rows && rows,
|
||||
rotate rot)
|
||||
{
|
||||
constexpr std::size_t n = detail_walk::domain_size<DpfKey>();
|
||||
detail_walk::rotated_weights<std::remove_reference_t<Rows>> view{
|
||||
rows, rot.shift % n, n};
|
||||
return eval_full_inner_product<I, Is...>(paired, dpf, view);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// eval_full_add_into(buf, key [, rotate | sketch])
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// @brief Add a full-domain expansion of output `I` into `buf` in place.
|
||||
/// @details `buf[i] += DPF_I(i)` (the leaf share's group element) for every
|
||||
/// domain point `i`. `buf` already holds the caller's running shares
|
||||
/// (Prio's histogram, Express's mailbox, Duoram's update); its element
|
||||
/// type must support `+` with the leaf share's group element.
|
||||
/// \complexity One full-domain expansion, `Θ(2^n)`, plus one `+` per point.
|
||||
template <std::size_t I = 0, typename Buffer, typename DpfKey,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
void eval_full_add_into(Buffer & buf, const DpfKey & dpf) // NOLINT(runtime/references)
|
||||
{
|
||||
auto result = eval_full<I>(dpf);
|
||||
auto & iter = result.second;
|
||||
std::size_t i = 0;
|
||||
for (auto it = std::begin(iter); it != std::end(iter); ++it, ++i)
|
||||
buf[i] = buf[i] + detail_walk::group_value(*it);
|
||||
}
|
||||
|
||||
/// @brief Add a full-domain expansion into `buf`, shifted by `rot`.
|
||||
/// @details `buf[(i + rot.shift) mod 2^n] += DPF_I(i)`. Duoram's update writes
|
||||
/// the payload placed at `r` into the memory slot `r + shift`.
|
||||
/// \complexity One full-domain expansion, `Θ(2^n)`, plus one `+` per point.
|
||||
template <std::size_t I = 0, typename Buffer, typename DpfKey,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
void eval_full_add_into(Buffer & buf, const DpfKey & dpf, rotate rot) // NOLINT(runtime/references)
|
||||
{
|
||||
constexpr std::size_t n = detail_walk::domain_size<DpfKey>();
|
||||
const std::size_t s = rot.shift % n;
|
||||
auto result = eval_full<I>(dpf);
|
||||
auto & iter = result.second;
|
||||
std::size_t i = 0;
|
||||
for (auto it = std::begin(iter); it != std::end(iter); ++it, ++i)
|
||||
{
|
||||
const std::size_t j = (i + s) % n;
|
||||
buf[j] = buf[j] + detail_walk::group_value(*it);
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Add a full-domain expansion into `buf` and fold the audit sketch.
|
||||
/// @details `buf[i] += DPF_I(i)` for every point, and each written share is
|
||||
/// absorbed into `sk` — the same one-hot audit as
|
||||
/// `eval_full(key, sketch(σ))`, in the same pass. Express's mailbox
|
||||
/// write becomes one call instead of a point loop and a second fold.
|
||||
/// \complexity One full-domain expansion, `Θ(2^n)`, plus one `+` and one
|
||||
/// absorb per point.
|
||||
template <std::size_t I = 0, typename Buffer, typename DpfKey,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
void eval_full_add_into(Buffer & buf, const DpfKey & dpf, sketch_ref sk) // NOLINT(runtime/references)
|
||||
{
|
||||
static_assert(DpfKey::is_extractable,
|
||||
"eval_full_add_into(..., sketch(σ)): key must carry dpf::extractable");
|
||||
auto result = eval_full<I>(dpf);
|
||||
auto & iter = result.second;
|
||||
std::size_t i = 0;
|
||||
for (auto it = std::begin(iter); it != std::end(iter); ++it, ++i)
|
||||
{
|
||||
auto g = detail_walk::group_value(*it);
|
||||
buf[i] = buf[i] + g;
|
||||
sk.absorb(g); // fold the raw share, as eval_point(..., sketch) does
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// dpf::cyclic_shift(buffer, s)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// @brief A copy of `buf` rotated so `result[i] == buf[(i - s) mod n]`.
|
||||
/// @details The buffer analogue of `dpf::rotate{s}`: the value at index `k`
|
||||
/// moves to `k + s`. Duoram's update shifts an expanded share buffer
|
||||
/// this way before adding it into memory.
|
||||
template <typename Container>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
Container cyclic_shift(const Container & buf, std::size_t s)
|
||||
{
|
||||
Container out(buf);
|
||||
const std::size_t n = buf.size();
|
||||
if (n == 0)
|
||||
return out;
|
||||
const std::size_t sh = s % n;
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
out[i] = buf[(i + n - sh) % n];
|
||||
return out;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// dpf::pack_bit_columns(keys...)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// @brief Pack the full-domain bit expansions of `keys` into one integer per row.
|
||||
/// @details Runs `eval_full` once per key over a shared bit domain and sets bit
|
||||
/// `e` of `result[row]` when key `e` opens to 1 at that row. `L` keys
|
||||
/// need `L <= 8*sizeof(Int)`. BitMore reads `result[row]` as the
|
||||
/// server's `L`-bit digit instead of unpacking one `int` per bit.
|
||||
/// @tparam Int packed digit type (defaults to `std::uint64_t`)
|
||||
/// \complexity One full-domain bit expansion per key, `Θ(L · 2^n)`.
|
||||
template <typename Int = std::uint64_t, typename First, typename... Rest>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::vector<Int> pack_bit_columns(const First & first, const Rest &... rest)
|
||||
{
|
||||
static_assert(1 + sizeof...(Rest) <= 8 * sizeof(Int),
|
||||
"pack_bit_columns: more keys than bits in the packed digit type");
|
||||
constexpr std::size_t n = detail_walk::domain_size<First>();
|
||||
std::vector<Int> out(n, Int{0});
|
||||
|
||||
std::size_t e = 0;
|
||||
auto do_one = [&](const auto & key)
|
||||
{
|
||||
auto result = eval_full(key);
|
||||
auto & iter = result.second;
|
||||
std::size_t row = 0;
|
||||
for (auto it = std::begin(iter); it != std::end(iter) && row < n;
|
||||
++it, ++row)
|
||||
{
|
||||
if (static_cast<bool>(*it))
|
||||
out[row] |= static_cast<Int>(Int{1} << e);
|
||||
}
|
||||
++e;
|
||||
};
|
||||
do_one(first);
|
||||
(do_one(rest), ...);
|
||||
return out;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// dpf::mod_bit_columns<ℓ>(keys...)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// @brief One residue per row: the packed bit columns, modulo `Modulus`.
|
||||
/// @details The same full-domain bit walk as `pack_bit_columns`. Key 0 is the
|
||||
/// low bit, so row `r` opens to `(sum_e bit_e(r) · 2^e) mod Modulus`.
|
||||
/// Bits are folded high-bit first. Through modulus 128 each running
|
||||
/// slot is a byte; through 32768 it is a 16-bit lane. Both use the
|
||||
/// high-nibble partial reduction, and only the finished row is fully
|
||||
/// reduced. Hafiz and Henry §5.3 read `result[row]` as the server's
|
||||
/// digit when the server count is not a power of two.
|
||||
/// @tparam Modulus server count, `2` through `32768`
|
||||
/// @tparam Int residue type (defaults to `std::uint16_t`)
|
||||
/// \complexity One full-domain bit expansion per key, `Θ(L · 2^n)`, and one
|
||||
/// lane insertion per row per key.
|
||||
template <unsigned Modulus, typename Int = std::uint16_t, typename First, typename... Rest>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::vector<Int> mod_bit_columns(const First & first, const Rest &... rest)
|
||||
{
|
||||
static_assert(Modulus >= 2u && Modulus <= 32768u,
|
||||
"mod_bit_columns: modulus must be in 2..32768");
|
||||
static_assert(std::is_unsigned_v<Int>,
|
||||
"mod_bit_columns: residue type must be unsigned");
|
||||
static_assert(Modulus - 1u <= static_cast<std::uintmax_t>(std::numeric_limits<Int>::max()),
|
||||
"mod_bit_columns: residue type cannot hold a value modulo Modulus");
|
||||
|
||||
constexpr unsigned lane_bits = Modulus <= 128u ? 8u : 16u;
|
||||
using reg = simde__m256i;
|
||||
using acc = bitmore_accumulator<Modulus, reg, lane_bits>;
|
||||
constexpr std::size_t lanes = lane_bits == 8u ? 32u : 16u;
|
||||
constexpr std::size_t n = detail_walk::domain_size<First>();
|
||||
const std::size_t chunks = (n + lanes - 1u) / lanes;
|
||||
|
||||
std::vector<acc> running(chunks);
|
||||
const auto keys = std::forward_as_tuple(first, rest...);
|
||||
detail_walk::insert_keys_msb_first([&](const auto & key)
|
||||
{
|
||||
using key_t = std::decay_t<decltype(key)>;
|
||||
static_assert(detail_walk::domain_size<key_t>() == n,
|
||||
"mod_bit_columns: every key must share the first key's domain");
|
||||
auto result = eval_full(key);
|
||||
auto & iter = result.second;
|
||||
auto it = std::begin(iter);
|
||||
const auto end = std::end(iter);
|
||||
for (std::size_t c = 0; c < chunks; ++c)
|
||||
{
|
||||
reg bits{};
|
||||
if constexpr (lane_bits == 8u)
|
||||
{
|
||||
alignas(32) unsigned char raw[32]{};
|
||||
for (std::size_t i = 0; i < lanes && it != end; ++i, ++it)
|
||||
raw[i] = static_cast<bool>(*it) ? 1u : 0u;
|
||||
std::memcpy(&bits, raw, sizeof(bits));
|
||||
}
|
||||
else
|
||||
{
|
||||
alignas(32) std::uint16_t raw[16]{};
|
||||
for (std::size_t i = 0; i < lanes && it != end; ++i, ++it)
|
||||
raw[i] = static_cast<bool>(*it) ? 1u : 0u;
|
||||
std::memcpy(&bits, raw, sizeof(bits));
|
||||
}
|
||||
running[c].insert_bit(bits);
|
||||
}
|
||||
}, keys, std::make_index_sequence<1u + sizeof...(Rest)>{});
|
||||
|
||||
std::vector<Int> out(n);
|
||||
std::size_t row = 0;
|
||||
for (std::size_t c = 0; c < chunks; ++c)
|
||||
{
|
||||
const reg reduced = running[c].reduced();
|
||||
if constexpr (lane_bits == 8u)
|
||||
{
|
||||
alignas(32) unsigned char raw[32]{};
|
||||
std::memcpy(raw, &reduced, sizeof(raw));
|
||||
for (std::size_t i = 0; i < lanes && row < n; ++i, ++row)
|
||||
out[row] = static_cast<Int>(raw[i]);
|
||||
}
|
||||
else
|
||||
{
|
||||
alignas(32) std::uint16_t raw[16]{};
|
||||
std::memcpy(raw, &reduced, sizeof(raw));
|
||||
for (std::size_t i = 0; i < lanes && row < n; ++i, ++row)
|
||||
out[row] = static_cast<Int>(raw[i]);
|
||||
}
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// eval_prefixes / eval_prefix_inner_product (idpf prefix walk)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// @brief The `2^N` prefix shares of output `I` (prefix length `N`).
|
||||
/// @details One walk to depth `N`; slot `p` is the share on prefix `p`. Poplar
|
||||
/// reads these to score every node at a depth in one pass instead of
|
||||
/// one `eval_point` per node. Returns the `(buffer, iterable)` pair of
|
||||
/// `eval_full(out<I,N>, key)`.
|
||||
/// \complexity One walk to depth `N`: `Θ(N + 2^N)` interior traversals.
|
||||
template <std::size_t I, std::size_t N, typename KeyT,
|
||||
std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_prefixes(out_t<I, N>, const KeyT & key)
|
||||
{
|
||||
return eval_full(out_t<I, N>{}, key);
|
||||
}
|
||||
|
||||
/// @brief `sum_p DPF_I(p) * values[p]` over the `2^N` prefixes of output `I`.
|
||||
/// @details Walks to depth `N` once and dots the prefix shares with `values`
|
||||
/// (indexed by prefix `0 .. 2^N - 1`). PRAC's strides and Poplar's
|
||||
/// "is this prefix heavy?" are this one call, replacing an
|
||||
/// `eval_point` per node.
|
||||
/// \complexity One walk to depth `N`, `Θ(N + 2^N)`, plus one multiply-add per
|
||||
/// prefix into an `O(1)` accumulator.
|
||||
template <std::size_t I, std::size_t N, typename KeyT, typename Values,
|
||||
std::enable_if_t<is_multilevel_key_v<KeyT>, bool> = true>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_prefix_inner_product(out_t<I, N>, const KeyT & key, Values && values)
|
||||
{
|
||||
using lane_t = typename KeyT::input_type;
|
||||
constexpr lane_t lo = lane_t{0};
|
||||
constexpr lane_t hi = (N >= utils::bitlength_of_v<lane_t>)
|
||||
? static_cast<lane_t>(~lane_t{0})
|
||||
: static_cast<lane_t>((lane_t{1} << N) - 1);
|
||||
return eval_inner_product(out_t<I, N>{}, key, lo, hi,
|
||||
std::forward<Values>(values));
|
||||
}
|
||||
|
||||
/// @brief XOR of `records[j]` for each listed point where the bit share is set.
|
||||
/// @details Same walk as `eval_sequence` on a bit key, but folds the selected
|
||||
/// records into one accumulator instead of materializing a bit vector.
|
||||
/// Keyword PIR's server response is this one call.
|
||||
template <typename DpfKey, typename ForwardIterator, typename Records,
|
||||
std::enable_if_t<!has_embedded_dpf_key_v<std::decay_t<DpfKey>>, int> = 0>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_sequence_xor(const DpfKey & key, ForwardIterator begin,
|
||||
ForwardIterator end, Records && records)
|
||||
{
|
||||
using record_t = std::decay_t<decltype(records[0])>;
|
||||
record_t acc{};
|
||||
auto [buf, iter] = eval_sequence(key, begin, end);
|
||||
std::size_t i = 0;
|
||||
for (auto bit : iter)
|
||||
{
|
||||
if (static_cast<bool>(bit))
|
||||
acc = static_cast<record_t>(acc ^ records[i]);
|
||||
++i;
|
||||
}
|
||||
return acc;
|
||||
}
|
||||
|
||||
/// @brief Evaluate a walk while invoking `fold(index, share)` once per written output.
|
||||
/// @details The fold is a template — inlined the way `sketch_ref::absorb` is.
|
||||
/// Express's audit, Pika's SZ check, and a SNIP input vector each
|
||||
/// supply their own fold over the group they write.
|
||||
template <typename Fold, typename Buffer, typename DpfKey,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>
|
||||
&& !std::is_same_v<std::decay_t<Fold>, sketch_ref>
|
||||
&& !std::is_same_v<std::decay_t<Fold>, rotate>, bool> = true>
|
||||
void eval_full_add_into(Buffer & buf, const DpfKey & dpf, Fold && fold) // NOLINT(runtime/references)
|
||||
{
|
||||
auto result = eval_full(dpf);
|
||||
auto & iter = result.second;
|
||||
std::size_t i = 0;
|
||||
for (auto it = std::begin(iter); it != std::end(iter); ++it, ++i)
|
||||
{
|
||||
auto g = detail_walk::group_value(*it);
|
||||
buf[i] = buf[i] + g;
|
||||
fold(i, g);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_EVAL_WALK_HPP__
|
||||
876
include/dpf/experiment.hpp
Normal file
876
include/dpf/experiment.hpp
Normal file
|
|
@ -0,0 +1,876 @@
|
|||
/// @file dpf/experiment.hpp
|
||||
/// @brief Replayable master seed + paper-ready protocol cost CSVs.
|
||||
/// @details A master covers the thread that installs it. `derive_party(i)`
|
||||
/// gives party `i` of a run its own stream, keyed by SHA-256 of the
|
||||
/// master and `i`, and `app::run_parties` installs one on each party
|
||||
/// thread, so replaying a master replays every party. Kernels handed
|
||||
/// to a compute pool draw from their party's stream
|
||||
/// (`dpf/thread_work.hpp`). Every CSV row carries the invocation id
|
||||
/// of the process that wrote it (`log::invocation_id()`), and a file
|
||||
/// whose header does not match this layout is moved aside rather
|
||||
/// than appended to.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_EXPERIMENT_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_EXPERIMENT_HPP__
|
||||
|
||||
#include <algorithm>
|
||||
#include <array>
|
||||
#include <chrono>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstdio>
|
||||
#include <cstdlib>
|
||||
#include <cstring>
|
||||
#include <filesystem>
|
||||
#include <fstream>
|
||||
#include <iomanip>
|
||||
#include <map>
|
||||
#include <sstream>
|
||||
#include <stdexcept>
|
||||
#include <string>
|
||||
#include <system_error>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include <time.h>
|
||||
|
||||
#include "dpf/compose.hpp"
|
||||
#include "dpf/experiment_note.hpp"
|
||||
#include "dpf/log.hpp"
|
||||
#include "dpf/prg_aes.hpp"
|
||||
#include "dpf/prg_count.hpp"
|
||||
#include "dpf/protocol.hpp"
|
||||
#include "dpf/random.hpp"
|
||||
#include "dpf/thread_work.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
using protocol::round_event;
|
||||
using protocol::round_probe;
|
||||
|
||||
/// @brief Per-thread experiment: master seed stream + cost meter + CSV export.
|
||||
class experiment : public detail::experiment_seed_sink
|
||||
{
|
||||
public:
|
||||
static constexpr std::size_t master_bytes = 32;
|
||||
|
||||
using master_seed = std::array<std::uint8_t, master_bytes>;
|
||||
|
||||
/// @brief Where a master came from.
|
||||
enum class origin : unsigned char
|
||||
{
|
||||
fresh, ///< drawn from OS entropy
|
||||
provided, ///< given to `replay`
|
||||
derived ///< `derive_party` of another master
|
||||
};
|
||||
|
||||
struct noted_seed
|
||||
{
|
||||
std::string name;
|
||||
std::vector<std::uint8_t> bytes;
|
||||
};
|
||||
|
||||
/// @brief Fresh master seed from the system entropy source.
|
||||
/// @details Always reads the OS entropy path (never an outer experiment
|
||||
/// hook), so nested contexts do not steal the parent stream.
|
||||
explicit experiment(std::string name, std::string party = "p0")
|
||||
: name_(std::move(name)), party_(std::move(party))
|
||||
{
|
||||
draw_system_master_();
|
||||
// Master itself must not count as protocol random consumption.
|
||||
reset_random_bytes_count();
|
||||
prg::reset_eval_count();
|
||||
install_();
|
||||
note("master", master_.data(), master_.size());
|
||||
log_master_();
|
||||
}
|
||||
|
||||
/// @brief Replay a previously recorded master seed.
|
||||
static experiment replay(std::string name, const master_seed & seed,
|
||||
std::string party = "p0")
|
||||
{
|
||||
return experiment(std::move(name), std::move(party), seed, origin::provided);
|
||||
}
|
||||
|
||||
/// @brief The stream party `party` of a run draws from, installed on the
|
||||
/// calling thread (call it on that party's thread). Its master is
|
||||
/// SHA-256 of this master and the party index, so one master
|
||||
/// always yields the same party streams.
|
||||
experiment derive_party(unsigned party) const
|
||||
{
|
||||
std::uint8_t in[master_bytes + 4];
|
||||
std::memcpy(in, master_.data(), master_bytes);
|
||||
for (int i = 0; i < 4; ++i)
|
||||
in[master_bytes + static_cast<std::size_t>(i)] =
|
||||
static_cast<std::uint8_t>((party >> (8 * i)) & 0xffu);
|
||||
const auto digest = log::detail::sha256("libdpf experiment party", in, sizeof(in));
|
||||
master_seed derived{};
|
||||
std::memcpy(derived.data(), digest.data(), derived.size());
|
||||
return experiment(name_, "p" + std::to_string(party), derived, origin::derived);
|
||||
}
|
||||
|
||||
experiment(const experiment &) = delete;
|
||||
experiment & operator=(const experiment &) = delete;
|
||||
|
||||
experiment(experiment && other) noexcept
|
||||
: name_(std::move(other.name_)),
|
||||
party_(std::move(other.party_)),
|
||||
run_id_(other.run_id_),
|
||||
master_(other.master_),
|
||||
origin_(other.origin_),
|
||||
aes_seed_(other.aes_seed_),
|
||||
ctr_(other.ctr_),
|
||||
buf_pos_(other.buf_pos_),
|
||||
prev_hook_(other.prev_hook_),
|
||||
prev_ctx_(other.prev_ctx_),
|
||||
prev_sink_(other.prev_sink_),
|
||||
installed_(other.installed_),
|
||||
seeds_(std::move(other.seeds_)),
|
||||
rounds_(std::move(other.rounds_)),
|
||||
edge_in_(std::move(other.edge_in_)),
|
||||
edge_out_(std::move(other.edge_out_)),
|
||||
edge_plan_out_(std::move(other.edge_plan_out_)),
|
||||
interactive_rounds_(other.interactive_rounds_),
|
||||
dag_depth_(other.dag_depth_),
|
||||
critical_path_(std::move(other.critical_path_)),
|
||||
wall_ns_(other.wall_ns_),
|
||||
cpu_ns_(other.cpu_ns_),
|
||||
prg_evals_(other.prg_evals_),
|
||||
sym_(other.sym_),
|
||||
random_bytes_(other.random_bytes_),
|
||||
bytes_in_(other.bytes_in_),
|
||||
bytes_out_(other.bytes_out_),
|
||||
plan_bytes_out_(other.plan_bytes_out_),
|
||||
config_(std::move(other.config_)),
|
||||
trials_(std::move(other.trials_)),
|
||||
party_trials_(std::move(other.party_trials_)),
|
||||
wire_(other.wire_),
|
||||
timing_started_(other.timing_started_),
|
||||
wall0_(other.wall0_),
|
||||
cpu0_(other.cpu0_)
|
||||
{
|
||||
std::memcpy(buf_, other.buf_, sizeof(buf_));
|
||||
other.installed_ = false;
|
||||
if (installed_)
|
||||
{
|
||||
detail::uniform_bytes_hook = &experiment::hook_;
|
||||
detail::uniform_bytes_ctx = this;
|
||||
detail::experiment_seed_sink_tls = this;
|
||||
}
|
||||
}
|
||||
|
||||
~experiment() override { uninstall_(); }
|
||||
|
||||
const std::string & name() const noexcept { return name_; }
|
||||
const std::string & party() const noexcept { return party_; }
|
||||
const master_seed & seed() const noexcept { return master_; }
|
||||
std::uint64_t run_id() const noexcept { return run_id_; }
|
||||
|
||||
void set_run_id(std::uint64_t id) noexcept { run_id_ = id; }
|
||||
void set_party(std::string party) { party_ = std::move(party); }
|
||||
|
||||
/// @brief Public alias of `note` for call-site labels.
|
||||
void note_seed(const char * name, const void * bytes, std::size_t n)
|
||||
{
|
||||
note(name, static_cast<const std::uint8_t *>(bytes), n);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
void note_seed(const char * name, const T & seed)
|
||||
{
|
||||
note_seed(name, &seed, sizeof(seed));
|
||||
}
|
||||
|
||||
void note(const char * name, const std::uint8_t * bytes,
|
||||
std::size_t n) override
|
||||
{
|
||||
if (name == nullptr || bytes == nullptr || n == 0)
|
||||
return;
|
||||
seeds_.push_back({name, std::vector<std::uint8_t>(bytes, bytes + n)});
|
||||
if (std::strcmp(name, "master") != 0)
|
||||
DPF_LOG(debug, "seed").kv("name", name).kv("experiment", name_)
|
||||
.kv("party", party_).kv("source", "noted").kv("bytes", n)
|
||||
.seed("value", bytes, n);
|
||||
}
|
||||
|
||||
/// @brief Noted seeds in order, starting with `master`.
|
||||
const std::vector<noted_seed> & seeds() const noexcept { return seeds_; }
|
||||
|
||||
/// @brief Append a party stream's noted seeds, each name prefixed by
|
||||
/// `party` (`p1/master`, `p1/buffered_prg`).
|
||||
void fold_seeds(const std::string & party, const std::vector<noted_seed> & seeds)
|
||||
{
|
||||
for (const auto & s : seeds)
|
||||
seeds_.push_back({party + "/" + s.name, s.bytes});
|
||||
}
|
||||
|
||||
origin seed_origin() const noexcept { return origin_; }
|
||||
|
||||
/// @brief True when the master was not drawn fresh (`replay` or
|
||||
/// `derive_party`).
|
||||
bool seed_provided() const noexcept { return origin_ != origin::fresh; }
|
||||
|
||||
static const char * origin_name(origin o) noexcept
|
||||
{
|
||||
switch (o)
|
||||
{
|
||||
case origin::fresh:
|
||||
return "fresh";
|
||||
case origin::provided:
|
||||
return "provided";
|
||||
case origin::derived:
|
||||
return "derived";
|
||||
}
|
||||
return "fresh";
|
||||
}
|
||||
|
||||
/// @brief Symmetric-key blocks counted between `begin_timing` and
|
||||
/// `end_timing`, by purpose and primitive (see `prg_count.hpp`).
|
||||
const prg::counts & sym_counts() const noexcept { return sym_; }
|
||||
|
||||
/// @brief Record static plan facts (chain length, per-edge schedule bytes).
|
||||
void ingest_plan(const protocol::plan & p)
|
||||
{
|
||||
interactive_rounds_ = p.rounds();
|
||||
dag_depth_ = p.waves();
|
||||
plan_bytes_out_ = 0;
|
||||
for (auto n : p.slot_bytes_all())
|
||||
plan_bytes_out_ += n;
|
||||
edge_plan_out_.clear();
|
||||
for (std::size_t wi = 0; wi < p.waves(); ++wi)
|
||||
{
|
||||
const auto & w = p.wave(wi);
|
||||
if (w.exchanges.empty())
|
||||
continue;
|
||||
const auto ch = protocol::detail::wave_channel(p, w);
|
||||
edge_plan_out_[static_cast<int>(ch)] += w.slot_bytes;
|
||||
}
|
||||
critical_path_ = build_critical_path_(p);
|
||||
}
|
||||
|
||||
/// @brief Probe that records live round events into this experiment.
|
||||
round_probe probe() noexcept
|
||||
{
|
||||
return round_probe{this, &experiment::on_round_};
|
||||
}
|
||||
|
||||
/// @brief Start wall/CPU/PRG/random timers (call before drive).
|
||||
void begin_timing()
|
||||
{
|
||||
prg::reset_eval_count();
|
||||
reset_random_bytes_count();
|
||||
wall0_ = steady_now_();
|
||||
cpu0_ = thread_cpu_now_();
|
||||
timing_started_ = true;
|
||||
}
|
||||
|
||||
/// @brief Stop timers and fold totals.
|
||||
void end_timing()
|
||||
{
|
||||
if (!timing_started_)
|
||||
return;
|
||||
wall_ns_ += steady_now_() - wall0_;
|
||||
cpu_ns_ += thread_cpu_now_() - cpu0_;
|
||||
prg_evals_ += prg::eval_count();
|
||||
const auto sym = prg::snapshot();
|
||||
for (std::size_t i = 0; i < sym.size(); ++i)
|
||||
sym_[i] += sym[i];
|
||||
random_bytes_ += random_bytes_count();
|
||||
timing_started_ = false;
|
||||
}
|
||||
|
||||
std::uint64_t wall_ns() const noexcept { return wall_ns_; }
|
||||
std::uint64_t cpu_ns() const noexcept { return cpu_ns_; }
|
||||
std::uint64_t prg_evals() const noexcept { return prg_evals_; }
|
||||
std::uint64_t random_bytes() const noexcept { return random_bytes_; }
|
||||
std::size_t bytes_in() const noexcept { return bytes_in_; }
|
||||
std::size_t bytes_out() const noexcept { return bytes_out_; }
|
||||
std::size_t interactive_rounds() const noexcept
|
||||
{
|
||||
return interactive_rounds_;
|
||||
}
|
||||
std::size_t dag_depth() const noexcept { return dag_depth_; }
|
||||
std::size_t plan_bytes_out() const noexcept { return plan_bytes_out_; }
|
||||
const std::vector<round_event> & rounds() const noexcept { return rounds_; }
|
||||
std::size_t seed_count() const noexcept { return seeds_.size(); }
|
||||
std::size_t critical_path_length() const noexcept
|
||||
{
|
||||
return critical_path_.size();
|
||||
}
|
||||
|
||||
/// @brief True if a seed named `name` was noted (including `"master"`).
|
||||
bool has_seed_named(const char * name) const
|
||||
{
|
||||
if (name == nullptr)
|
||||
return false;
|
||||
for (const auto & s : seeds_)
|
||||
if (s.name == name)
|
||||
return true;
|
||||
return false;
|
||||
}
|
||||
|
||||
/// @brief Schedule bytes attributed to `channel` by `ingest_plan`.
|
||||
std::size_t plan_edge_bytes(protocol::edge_channel channel) const
|
||||
{
|
||||
return edge_get_(edge_plan_out_, static_cast<int>(channel));
|
||||
}
|
||||
|
||||
/// @brief Live bytes in/out attributed to `channel` by the round probe.
|
||||
std::size_t edge_bytes_in(protocol::edge_channel channel) const
|
||||
{
|
||||
return edge_get_(edge_in_, static_cast<int>(channel));
|
||||
}
|
||||
std::size_t edge_bytes_out(protocol::edge_channel channel) const
|
||||
{
|
||||
return edge_get_(edge_out_, static_cast<int>(channel));
|
||||
}
|
||||
|
||||
std::string seed_hex() const { return to_hex_(master_.data(), master_.size()); }
|
||||
|
||||
/// @brief On-the-wire counters for party 0's links (headers included).
|
||||
struct wire_counts
|
||||
{
|
||||
std::uint64_t bytes_out = 0;
|
||||
std::uint64_t bytes_in = 0;
|
||||
std::uint64_t payload_out = 0;
|
||||
std::uint64_t payload_in = 0;
|
||||
std::uint64_t frames_out = 0;
|
||||
std::uint64_t frames_in = 0;
|
||||
std::uint64_t write_calls = 0;
|
||||
};
|
||||
|
||||
/// @brief Run configuration recorded in `config.csv` (key, value).
|
||||
void set_config(std::vector<std::pair<std::string, std::string>> kv)
|
||||
{
|
||||
config_ = std::move(kv);
|
||||
}
|
||||
const std::vector<std::pair<std::string, std::string>> & config() const noexcept
|
||||
{
|
||||
return config_;
|
||||
}
|
||||
|
||||
/// @brief One timed trial: party 0's wall time and, when given, every
|
||||
/// party's (`party_walls[i]` is party `i`). Recorded in `trials.csv`.
|
||||
void add_trial(std::uint64_t wall_ns,
|
||||
const std::vector<std::uint64_t> & party_walls = {})
|
||||
{
|
||||
trials_.push_back(wall_ns);
|
||||
party_trials_.push_back(party_walls);
|
||||
}
|
||||
const std::vector<std::uint64_t> & trials() const noexcept { return trials_; }
|
||||
|
||||
/// @brief Median of party 0's trial wall times (0 when none).
|
||||
std::uint64_t median_trial_ns() const { return median_(trials_); }
|
||||
|
||||
/// @brief Median over trials of the slowest party's wall time (party 0's
|
||||
/// when a trial did not record the others).
|
||||
std::uint64_t slowest_median_ns() const
|
||||
{
|
||||
std::vector<std::uint64_t> slowest;
|
||||
for (std::size_t t = 0; t < trials_.size(); ++t)
|
||||
{
|
||||
std::uint64_t w = trials_[t];
|
||||
if (t < party_trials_.size())
|
||||
for (auto p : party_trials_[t])
|
||||
w = std::max(w, p);
|
||||
slowest.push_back(w);
|
||||
}
|
||||
return median_(slowest);
|
||||
}
|
||||
|
||||
void set_wire(const wire_counts & w) { wire_ = w; }
|
||||
const wire_counts & wire() const noexcept { return wire_; }
|
||||
|
||||
/// @brief Write / append CSV tables under `dir` (created if missing).
|
||||
void write_csv(const std::string & dir) const
|
||||
{
|
||||
if (dir.empty())
|
||||
throw std::invalid_argument("experiment::write_csv empty dir");
|
||||
std::error_code ec;
|
||||
std::filesystem::create_directories(dir, ec);
|
||||
if (ec)
|
||||
throw std::runtime_error("experiment: cannot create '" + dir + "': "
|
||||
+ ec.message());
|
||||
write_summary_(dir);
|
||||
write_rounds_(dir);
|
||||
write_edges_(dir);
|
||||
write_seeds_(dir);
|
||||
write_critical_path_(dir);
|
||||
write_config_(dir);
|
||||
write_trials_(dir);
|
||||
write_wire_(dir);
|
||||
write_sym_(dir);
|
||||
write_runs_(dir);
|
||||
DPF_LOG(info, "csv").kv("dir", dir).kv("experiment", name_)
|
||||
.kv("party", party_).kv("run_id", run_id_).kv("rounds", rounds_.size())
|
||||
.kv("trials", trials_.size()).kv("seeds", seeds_.size());
|
||||
}
|
||||
|
||||
private:
|
||||
experiment(std::string name, std::string party, master_seed seed, origin o)
|
||||
: name_(std::move(name)), party_(std::move(party)), master_(seed), origin_(o)
|
||||
{
|
||||
reset_random_bytes_count();
|
||||
prg::reset_eval_count();
|
||||
install_();
|
||||
note("master", master_.data(), master_.size());
|
||||
log_master_();
|
||||
}
|
||||
|
||||
static const char * entropy_name_() noexcept
|
||||
{
|
||||
#if defined(LIBDPF_USE_ARC4RANDOM)
|
||||
return "arc4random";
|
||||
#elif defined(LIBDPF_USE_DEV_RANDOM)
|
||||
return "/dev/random";
|
||||
#else
|
||||
return "/dev/urandom";
|
||||
#endif
|
||||
}
|
||||
|
||||
void log_master_() const
|
||||
{
|
||||
const bool derived = origin_ == origin::derived;
|
||||
if (!log::enabled(derived ? log::level::debug : log::level::info))
|
||||
return;
|
||||
log::record(derived ? log::level::debug : log::level::info, "seed")
|
||||
.kv("name", "master").kv("experiment", name_).kv("party", party_)
|
||||
.kv("source", origin_name(origin_))
|
||||
.kv("entropy", origin_ == origin::fresh ? entropy_name_()
|
||||
: derived ? "sha256(master,party)"
|
||||
: "replay")
|
||||
.kv("bytes", master_.size()).seed("value", master_.data(), master_.size());
|
||||
}
|
||||
|
||||
void draw_system_master_()
|
||||
{
|
||||
// Bypass any installed hook so the master is true OS entropy.
|
||||
auto * saved = detail::uniform_bytes_hook;
|
||||
detail::uniform_bytes_hook = nullptr;
|
||||
for (auto & b : master_)
|
||||
uniform_fill(b);
|
||||
detail::uniform_bytes_hook = saved;
|
||||
}
|
||||
|
||||
void install_()
|
||||
{
|
||||
prev_hook_ = detail::uniform_bytes_hook;
|
||||
prev_ctx_ = detail::uniform_bytes_ctx;
|
||||
prev_sink_ = detail::experiment_seed_sink_tls;
|
||||
detail::uniform_bytes_hook = &experiment::hook_;
|
||||
detail::uniform_bytes_ctx = this;
|
||||
detail::experiment_seed_sink_tls = this;
|
||||
installed_ = true;
|
||||
// Derive AES key from the first 16 master bytes.
|
||||
std::memcpy(&aes_seed_, master_.data(), sizeof(aes_seed_));
|
||||
ctr_ = 0;
|
||||
buf_pos_ = sizeof(buf_); // force refill
|
||||
}
|
||||
|
||||
void uninstall_()
|
||||
{
|
||||
if (!installed_)
|
||||
return;
|
||||
if (detail::uniform_bytes_ctx == this)
|
||||
detail::uniform_bytes_ctx = prev_ctx_;
|
||||
if (detail::uniform_bytes_hook == &experiment::hook_)
|
||||
detail::uniform_bytes_hook = prev_hook_;
|
||||
if (detail::experiment_seed_sink_tls == this)
|
||||
detail::experiment_seed_sink_tls = prev_sink_;
|
||||
installed_ = false;
|
||||
}
|
||||
|
||||
static void hook_(void * dst, std::size_t n)
|
||||
{
|
||||
auto * self = static_cast<experiment *>(detail::uniform_bytes_ctx);
|
||||
if (self == nullptr)
|
||||
throw std::logic_error("experiment hook without active context");
|
||||
self->fill_(dst, n);
|
||||
}
|
||||
|
||||
void fill_(void * dst, std::size_t n)
|
||||
{
|
||||
auto * out = static_cast<std::uint8_t *>(dst);
|
||||
while (n > 0)
|
||||
{
|
||||
if (buf_pos_ >= sizeof(buf_))
|
||||
{
|
||||
const prg::purpose_scope harness(prg::purpose::harness);
|
||||
auto blk = prg::aes128::eval(aes_seed_,
|
||||
static_cast<psnip_uint32_t>(ctr_++));
|
||||
std::memcpy(buf_, &blk, sizeof(buf_));
|
||||
buf_pos_ = 0;
|
||||
}
|
||||
const std::size_t take =
|
||||
std::min(n, sizeof(buf_) - buf_pos_);
|
||||
std::memcpy(out, buf_ + buf_pos_, take);
|
||||
buf_pos_ += take;
|
||||
out += take;
|
||||
n -= take;
|
||||
}
|
||||
}
|
||||
|
||||
static void on_round_(void * ctx, const round_event & ev)
|
||||
{
|
||||
static_cast<experiment *>(ctx)->record_round_(ev);
|
||||
}
|
||||
|
||||
void record_round_(const round_event & ev)
|
||||
{
|
||||
rounds_.push_back(ev);
|
||||
bytes_in_ += ev.bytes_in;
|
||||
bytes_out_ += ev.bytes_out;
|
||||
edge_in_[static_cast<int>(ev.channel)] += ev.bytes_in;
|
||||
edge_out_[static_cast<int>(ev.channel)] += ev.bytes_out;
|
||||
// wall / cpu / prg / random totals come from begin_timing/end_timing
|
||||
// so finish_schedule local work is included once.
|
||||
}
|
||||
|
||||
static std::uint64_t steady_now_()
|
||||
{
|
||||
using clock = std::chrono::steady_clock;
|
||||
return static_cast<std::uint64_t>(
|
||||
std::chrono::duration_cast<std::chrono::nanoseconds>(
|
||||
clock::now().time_since_epoch())
|
||||
.count());
|
||||
}
|
||||
|
||||
static std::uint64_t thread_cpu_now_()
|
||||
{
|
||||
if (have_thread_cpu_clock())
|
||||
return thread_cpu_ns();
|
||||
return steady_now_();
|
||||
}
|
||||
|
||||
static std::uint64_t median_(std::vector<std::uint64_t> v)
|
||||
{
|
||||
if (v.empty())
|
||||
return 0;
|
||||
std::sort(v.begin(), v.end());
|
||||
return v[v.size() / 2];
|
||||
}
|
||||
|
||||
static std::string to_hex_(const std::uint8_t * p, std::size_t n)
|
||||
{
|
||||
std::ostringstream os;
|
||||
os << std::hex << std::setfill('0');
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
os << std::setw(2) << static_cast<unsigned>(p[i]);
|
||||
return os.str();
|
||||
}
|
||||
|
||||
static const char * channel_name_(protocol::edge_channel c)
|
||||
{
|
||||
switch (c)
|
||||
{
|
||||
case protocol::edge_channel::peer:
|
||||
return "peer";
|
||||
case protocol::edge_channel::rss_next:
|
||||
return "rss_next";
|
||||
case protocol::edge_channel::dealer:
|
||||
return "dealer";
|
||||
}
|
||||
return "edge";
|
||||
}
|
||||
|
||||
struct path_node
|
||||
{
|
||||
std::uint32_t id = 0;
|
||||
std::size_t wave = 0;
|
||||
std::uint32_t opcode = 0;
|
||||
int effect = 0;
|
||||
};
|
||||
|
||||
static std::vector<path_node> build_critical_path_(const protocol::plan & p)
|
||||
{
|
||||
std::vector<path_node> path;
|
||||
if (p.nodes().empty())
|
||||
return path;
|
||||
std::uint32_t tip = p.nodes().front().id;
|
||||
std::size_t best_w = p.wave_of(protocol::node{tip});
|
||||
for (auto n : p.nodes())
|
||||
{
|
||||
const auto w = p.wave_of(n);
|
||||
if (w >= best_w)
|
||||
{
|
||||
best_w = w;
|
||||
tip = n.id;
|
||||
}
|
||||
}
|
||||
for (;;)
|
||||
{
|
||||
path_node pn;
|
||||
pn.id = tip;
|
||||
pn.wave = p.wave_of(protocol::node{tip});
|
||||
pn.opcode = p.opcode_of(tip);
|
||||
pn.effect = static_cast<int>(p.effect_of(tip));
|
||||
path.push_back(pn);
|
||||
const auto & ins = p.inputs_of(tip);
|
||||
if (ins.empty())
|
||||
break;
|
||||
std::uint32_t next = ins.front();
|
||||
std::size_t nw = p.wave_of(protocol::node{next});
|
||||
for (auto in : ins)
|
||||
{
|
||||
const auto w = p.wave_of(protocol::node{in});
|
||||
if (w >= nw)
|
||||
{
|
||||
nw = w;
|
||||
next = in;
|
||||
}
|
||||
}
|
||||
if (next == tip)
|
||||
break;
|
||||
tip = next;
|
||||
}
|
||||
std::reverse(path.begin(), path.end());
|
||||
return path;
|
||||
}
|
||||
|
||||
/// @brief Open `dir/file` for append, writing `header` first when the file
|
||||
/// is new. A file with a different header is renamed to
|
||||
/// `<stem>.before-<UTC>.csv` so rows never land under the wrong
|
||||
/// columns.
|
||||
static std::ofstream open_table_(const std::string & dir, const char * file,
|
||||
const char * header)
|
||||
{
|
||||
const std::string path = dir + "/" + file;
|
||||
std::string existing;
|
||||
{
|
||||
std::ifstream in(path);
|
||||
if (in)
|
||||
std::getline(in, existing);
|
||||
}
|
||||
if (!existing.empty() && existing != header)
|
||||
{
|
||||
std::string stamp;
|
||||
for (char c : log::detail::utc_text(std::chrono::system_clock::now()))
|
||||
if (c != '-' && c != ':')
|
||||
stamp += c;
|
||||
const std::string moved = path.substr(0, path.size() - 4) + ".before-"
|
||||
+ stamp + ".csv";
|
||||
if (std::rename(path.c_str(), moved.c_str()) != 0)
|
||||
throw std::runtime_error("experiment: " + path
|
||||
+ " has another column layout and cannot be moved aside");
|
||||
DPF_LOG(warning, "csv.moved").kv("file", path).kv("to", moved)
|
||||
.kv("detail", "its header differs from this build's columns");
|
||||
existing.clear();
|
||||
}
|
||||
std::ofstream out(path, std::ios::app);
|
||||
if (!out)
|
||||
throw std::runtime_error(std::string("experiment: cannot write ") + file);
|
||||
if (existing.empty())
|
||||
out << header << '\n';
|
||||
return out;
|
||||
}
|
||||
|
||||
void write_summary_(const std::string & dir) const
|
||||
{
|
||||
auto out = open_table_(dir, "summary.csv",
|
||||
"name,party,run_id,master_seed,interactive_rounds,dag_depth,"
|
||||
"wall_ns,cpu_ns,prg_evals,random_bytes,bytes_in,bytes_out,"
|
||||
"plan_bytes_out,peer_in,peer_out,rss_in,rss_out,dealer_in,"
|
||||
"dealer_out,median_ns,slowest_median_ns,trials,invocation");
|
||||
out << csv_escape_(name_) << ',' << csv_escape_(party_) << ','
|
||||
<< run_id_ << ',' << seed_hex() << ',' << interactive_rounds_ << ','
|
||||
<< dag_depth_ << ',' << wall_ns_ << ',' << cpu_ns_ << ','
|
||||
<< prg_evals_ << ',' << random_bytes_ << ',' << bytes_in_ << ','
|
||||
<< bytes_out_ << ',' << plan_bytes_out_ << ','
|
||||
<< edge_get_(edge_in_, 0) << ',' << edge_get_(edge_out_, 0) << ','
|
||||
<< edge_get_(edge_in_, 1) << ',' << edge_get_(edge_out_, 1) << ','
|
||||
<< edge_get_(edge_in_, 2) << ',' << edge_get_(edge_out_, 2) << ','
|
||||
<< median_trial_ns() << ',' << slowest_median_ns() << ','
|
||||
<< trials_.size() << ',' << log::invocation_id() << '\n';
|
||||
}
|
||||
|
||||
void write_rounds_(const std::string & dir) const
|
||||
{
|
||||
auto out = open_table_(dir, "rounds.csv",
|
||||
"name,party,run_id,round,edge,channel,bytes_out,bytes_in,"
|
||||
"wall_ns,cpu_ns,prg_evals,random_bytes,invocation");
|
||||
for (const auto & r : rounds_)
|
||||
{
|
||||
out << csv_escape_(name_) << ',' << csv_escape_(party_) << ','
|
||||
<< run_id_ << ',' << r.round << ',' << r.edge << ','
|
||||
<< channel_name_(r.channel) << ',' << r.bytes_out << ','
|
||||
<< r.bytes_in << ',' << r.wall_ns << ',' << r.cpu_ns << ','
|
||||
<< r.prg_evals << ',' << r.random_bytes << ','
|
||||
<< log::invocation_id() << '\n';
|
||||
}
|
||||
}
|
||||
|
||||
void write_edges_(const std::string & dir) const
|
||||
{
|
||||
auto out = open_table_(dir, "edges.csv",
|
||||
"name,party,run_id,channel,bytes_in,bytes_out,plan_bytes_out,invocation");
|
||||
for (int c = 0; c < 3; ++c)
|
||||
{
|
||||
const auto ch = static_cast<protocol::edge_channel>(c);
|
||||
out << csv_escape_(name_) << ',' << csv_escape_(party_) << ','
|
||||
<< run_id_ << ',' << channel_name_(ch) << ','
|
||||
<< edge_get_(edge_in_, c) << ',' << edge_get_(edge_out_, c)
|
||||
<< ',' << edge_get_(edge_plan_out_, c) << ','
|
||||
<< log::invocation_id() << '\n';
|
||||
}
|
||||
}
|
||||
|
||||
void write_seeds_(const std::string & dir) const
|
||||
{
|
||||
auto out = open_table_(dir, "seeds.csv",
|
||||
"name,party,run_id,seed_name,seed_hex,seed_bytes,invocation");
|
||||
for (const auto & s : seeds_)
|
||||
{
|
||||
out << csv_escape_(name_) << ',' << csv_escape_(party_) << ','
|
||||
<< run_id_ << ',' << csv_escape_(s.name) << ','
|
||||
<< to_hex_(s.bytes.data(), s.bytes.size()) << ','
|
||||
<< s.bytes.size() << ',' << log::invocation_id() << '\n';
|
||||
}
|
||||
}
|
||||
|
||||
void write_critical_path_(const std::string & dir) const
|
||||
{
|
||||
auto out = open_table_(dir, "critical_path.csv",
|
||||
"name,party,run_id,step,node_id,wave,opcode,effect,invocation");
|
||||
for (std::size_t i = 0; i < critical_path_.size(); ++i)
|
||||
{
|
||||
const auto & n = critical_path_[i];
|
||||
out << csv_escape_(name_) << ',' << csv_escape_(party_) << ','
|
||||
<< run_id_ << ',' << i << ',' << n.id << ',' << n.wave << ','
|
||||
<< n.opcode << ',' << n.effect << ',' << log::invocation_id() << '\n';
|
||||
}
|
||||
}
|
||||
|
||||
void write_config_(const std::string & dir) const
|
||||
{
|
||||
auto out = open_table_(dir, "config.csv", "name,party,run_id,key,value,invocation");
|
||||
for (const auto & kv : config_)
|
||||
out << csv_escape_(name_) << ',' << csv_escape_(party_) << ','
|
||||
<< run_id_ << ',' << csv_escape_(kv.first) << ','
|
||||
<< csv_escape_(kv.second) << ',' << log::invocation_id() << '\n';
|
||||
}
|
||||
|
||||
/// @brief One row per party per trial when parties were recorded, one row
|
||||
/// per trial for this experiment's party otherwise.
|
||||
void write_trials_(const std::string & dir) const
|
||||
{
|
||||
auto out = open_table_(dir, "trials.csv",
|
||||
"name,party,run_id,trial,wall_ns,invocation");
|
||||
for (std::size_t t = 0; t < trials_.size(); ++t)
|
||||
{
|
||||
const bool all = t < party_trials_.size() && !party_trials_[t].empty();
|
||||
if (!all)
|
||||
{
|
||||
out << csv_escape_(name_) << ',' << csv_escape_(party_) << ','
|
||||
<< run_id_ << ',' << t << ',' << trials_[t] << ','
|
||||
<< log::invocation_id() << '\n';
|
||||
continue;
|
||||
}
|
||||
for (std::size_t p = 0; p < party_trials_[t].size(); ++p)
|
||||
out << csv_escape_(name_) << ",p" << p << ',' << run_id_ << ','
|
||||
<< t << ',' << party_trials_[t][p] << ','
|
||||
<< log::invocation_id() << '\n';
|
||||
}
|
||||
}
|
||||
|
||||
void write_wire_(const std::string & dir) const
|
||||
{
|
||||
auto out = open_table_(dir, "wire.csv",
|
||||
"name,party,run_id,bytes_out,bytes_in,payload_out,payload_in,"
|
||||
"frames_out,frames_in,write_calls,invocation");
|
||||
out << csv_escape_(name_) << ',' << csv_escape_(party_) << ',' << run_id_
|
||||
<< ',' << wire_.bytes_out << ',' << wire_.bytes_in << ','
|
||||
<< wire_.payload_out << ',' << wire_.payload_in << ','
|
||||
<< wire_.frames_out << ',' << wire_.frames_in << ','
|
||||
<< wire_.write_calls << ',' << log::invocation_id() << '\n';
|
||||
}
|
||||
|
||||
void write_sym_(const std::string & dir) const
|
||||
{
|
||||
auto out = open_table_(dir, "sym.csv",
|
||||
"name,party,run_id,purpose,primitive,blocks,invocation");
|
||||
for (std::size_t u = 0; u < prg::purpose_count; ++u)
|
||||
for (std::size_t p = 0; p < prg::primitive_count; ++p)
|
||||
{
|
||||
const auto n = sym_[u * prg::primitive_count + p];
|
||||
if (n == 0)
|
||||
continue;
|
||||
out << csv_escape_(name_) << ',' << csv_escape_(party_) << ','
|
||||
<< run_id_ << ','
|
||||
<< prg::purpose_name(static_cast<prg::purpose>(u)) << ','
|
||||
<< prg::primitive_name(static_cast<prg::primitive>(p)) << ','
|
||||
<< n << ',' << log::invocation_id() << '\n';
|
||||
}
|
||||
}
|
||||
|
||||
void write_runs_(const std::string & dir) const
|
||||
{
|
||||
auto out = open_table_(dir, "runs.csv",
|
||||
"name,party,run_id,invocation,seed_source,written_utc");
|
||||
out << csv_escape_(name_) << ',' << csv_escape_(party_) << ',' << run_id_
|
||||
<< ',' << log::invocation_id() << ',' << origin_name(origin_) << ','
|
||||
<< log::detail::utc_text(std::chrono::system_clock::now()) << '\n';
|
||||
}
|
||||
|
||||
static std::string csv_escape_(const std::string & s)
|
||||
{
|
||||
if (s.find_first_of(",\"\n\r") == std::string::npos)
|
||||
return s;
|
||||
std::string out = "\"";
|
||||
for (char c : s)
|
||||
{
|
||||
if (c == '"')
|
||||
out += "\"\"";
|
||||
else
|
||||
out += c;
|
||||
}
|
||||
out += '"';
|
||||
return out;
|
||||
}
|
||||
|
||||
static std::size_t edge_get_(const std::map<int, std::size_t> & m, int k)
|
||||
{
|
||||
auto it = m.find(k);
|
||||
return it == m.end() ? 0 : it->second;
|
||||
}
|
||||
|
||||
std::string name_;
|
||||
std::string party_;
|
||||
std::uint64_t run_id_ = 0;
|
||||
master_seed master_{};
|
||||
origin origin_ = origin::fresh;
|
||||
prg::aes128::block_type aes_seed_{};
|
||||
std::uint32_t ctr_ = 0;
|
||||
std::uint8_t buf_[16]{};
|
||||
std::size_t buf_pos_ = 16;
|
||||
void (*prev_hook_)(void *, std::size_t) = nullptr;
|
||||
void * prev_ctx_ = nullptr;
|
||||
detail::experiment_seed_sink * prev_sink_ = nullptr;
|
||||
bool installed_ = false;
|
||||
std::vector<noted_seed> seeds_;
|
||||
std::vector<round_event> rounds_;
|
||||
std::map<int, std::size_t> edge_in_;
|
||||
std::map<int, std::size_t> edge_out_;
|
||||
std::map<int, std::size_t> edge_plan_out_;
|
||||
std::size_t interactive_rounds_ = 0;
|
||||
std::size_t dag_depth_ = 0;
|
||||
std::vector<path_node> critical_path_;
|
||||
std::uint64_t wall_ns_ = 0;
|
||||
std::uint64_t cpu_ns_ = 0;
|
||||
std::uint64_t prg_evals_ = 0;
|
||||
prg::counts sym_{};
|
||||
std::uint64_t random_bytes_ = 0;
|
||||
std::size_t bytes_in_ = 0;
|
||||
std::size_t bytes_out_ = 0;
|
||||
std::size_t plan_bytes_out_ = 0;
|
||||
std::vector<std::pair<std::string, std::string>> config_;
|
||||
std::vector<std::uint64_t> trials_;
|
||||
std::vector<std::vector<std::uint64_t>> party_trials_;
|
||||
wire_counts wire_{};
|
||||
bool timing_started_ = false;
|
||||
std::uint64_t wall0_ = 0;
|
||||
std::uint64_t cpu0_ = 0;
|
||||
};
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_EXPERIMENT_HPP__
|
||||
52
include/dpf/experiment_note.hpp
Normal file
52
include/dpf/experiment_note.hpp
Normal file
|
|
@ -0,0 +1,52 @@
|
|||
/// @file dpf/experiment_note.hpp
|
||||
/// @brief Lightweight seed registry hook (no Asio / PRG includes).
|
||||
#ifndef LIBDPF_INCLUDE_DPF_EXPERIMENT_NOTE_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_EXPERIMENT_NOTE_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <type_traits>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace detail
|
||||
{
|
||||
|
||||
/// @brief Optional sink installed by `experiment` (TLS).
|
||||
struct experiment_seed_sink
|
||||
{
|
||||
virtual ~experiment_seed_sink() = default;
|
||||
virtual void note(const char * name, const std::uint8_t * bytes,
|
||||
std::size_t n) = 0;
|
||||
};
|
||||
|
||||
inline thread_local experiment_seed_sink * experiment_seed_sink_tls = nullptr;
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/// @brief Record a named seed when an `experiment` is installed on this thread.
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void note_experiment_seed(const char * name, const void * bytes,
|
||||
std::size_t n) noexcept
|
||||
{
|
||||
if (detail::experiment_seed_sink_tls == nullptr || bytes == nullptr
|
||||
|| n == 0 || name == nullptr)
|
||||
return;
|
||||
detail::experiment_seed_sink_tls->note(name,
|
||||
static_cast<const std::uint8_t *>(bytes), n);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void note_experiment_seed(const char * name, const T & seed) noexcept
|
||||
{
|
||||
static_assert(std::is_trivially_copyable_v<T>,
|
||||
"note_experiment_seed requires a trivially copyable seed");
|
||||
note_experiment_seed(name, &seed, sizeof(seed));
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_EXPERIMENT_NOTE_HPP__
|
||||
482
include/dpf/factory_gadgets.hpp
Normal file
482
include/dpf/factory_gadgets.hpp
Normal file
|
|
@ -0,0 +1,482 @@
|
|||
/// @file dpf/factory_gadgets.hpp
|
||||
/// @brief Online MPC gadgets as `net::stream_array` round sequences.
|
||||
/// @details These mirror the limb logic in `mpc::circuit` but speak only
|
||||
/// `factory::detail::read_pod` / `exchange_pod` on dealer and peer
|
||||
/// arrays. Higher-level circuit ops compose from them:
|
||||
/// - `gt` — two `a2b_online` passes plus compare ANDs
|
||||
/// - `trunc_exact` — open `x-r` then a carry AND chain (as in A2B)
|
||||
/// - `mux` — Beaver mul on `(sel, a-b)` plus local add of `b`
|
||||
#ifndef LIBDPF_INCLUDE_DPF_FACTORY_GADGETS_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_FACTORY_GADGETS_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <stdexcept>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/circuit.hpp"
|
||||
#include "dpf/edabit.hpp"
|
||||
#include "dpf/gilboa.hpp"
|
||||
#include "dpf/net/round_sink.hpp"
|
||||
#include "dpf/net/stream_array.hpp"
|
||||
#include "dpf/net/stream_mesh.hpp"
|
||||
#include "dpf/ot_pack.hpp"
|
||||
#include "dpf/protocol_factory.hpp"
|
||||
#include "dpf/bit_inject.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace factory
|
||||
{
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
inline std::uint64_t limb_mask(std::uint16_t limb)
|
||||
{
|
||||
if (limb == 0 || limb > 8)
|
||||
throw std::invalid_argument("factory_gadgets limb");
|
||||
return limb >= 8 ? ~std::uint64_t{0}
|
||||
: ((std::uint64_t{1} << (8u * limb)) - 1u);
|
||||
}
|
||||
|
||||
inline void ring_view_to_limbs(const ring_triple_view & v, std::uint16_t limb,
|
||||
std::uint64_t & a, std::uint64_t & b, std::uint64_t & c)
|
||||
{
|
||||
a = b = c = 0;
|
||||
const std::uint16_t use = v.limb != 0 ? v.limb : limb;
|
||||
std::memcpy(&a, v.a, use);
|
||||
std::memcpy(&b, v.b, use);
|
||||
std::memcpy(&c, v.c, use);
|
||||
const auto mask = limb_mask(limb);
|
||||
a &= mask;
|
||||
b &= mask;
|
||||
c &= mask;
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/// @brief Sum-open one ring share on `peer[stream]` (additive declassify).
|
||||
inline std::uint64_t open_sum_online(unsigned /*party*/, std::uint64_t mine,
|
||||
net::stream_array & peer, std::size_t stream, std::uint16_t limb)
|
||||
{
|
||||
const auto mask = detail::limb_mask(limb);
|
||||
std::uint64_t theirs = 0;
|
||||
detail::exchange_pod(peer, stream, mine & mask, theirs);
|
||||
return (mine + theirs) & mask;
|
||||
}
|
||||
|
||||
/// @brief Beaver multiplication: triple on `dealer[dealer_stream]`, opens on
|
||||
/// `peer[peer_d_stream]` and `peer[peer_e_stream]`.
|
||||
inline std::uint64_t beaver_mul_online(unsigned party, std::uint64_t x,
|
||||
std::uint64_t y, net::stream_array & peer, net::stream_array & dealer,
|
||||
std::size_t dealer_stream = 0, std::size_t peer_d_stream = 0,
|
||||
std::size_t peer_e_stream = 1, std::uint16_t limb = 8)
|
||||
{
|
||||
const auto mask = detail::limb_mask(limb);
|
||||
const auto view =
|
||||
detail::read_pod<ring_triple_view>(dealer, dealer_stream);
|
||||
std::uint64_t a = 0, b = 0, c = 0;
|
||||
detail::ring_view_to_limbs(view, limb, a, b, c);
|
||||
const auto d = open_sum_online(party, (x - a) & mask, peer, peer_d_stream,
|
||||
limb);
|
||||
const auto e = open_sum_online(party, (y - b) & mask, peer, peer_e_stream,
|
||||
limb);
|
||||
std::uint64_t z = (c + d * b + e * a) & mask;
|
||||
if (party == 0)
|
||||
z = (z + d * e) & mask;
|
||||
return z;
|
||||
}
|
||||
|
||||
/// @brief GMW AND on XOR bit shares; blind from `dealer[dealer_stream]`.
|
||||
inline std::uint8_t gmw_and_online(unsigned party, std::uint8_t p,
|
||||
std::uint8_t q, net::stream_array & peer, net::stream_array & dealer,
|
||||
std::size_t dealer_stream = 0, std::size_t peer_stream = 0)
|
||||
{
|
||||
const auto blind = detail::read_pod<gmw_and_blind>(dealer, dealer_stream);
|
||||
gmw_and_msg mine{};
|
||||
mine.d = static_cast<std::uint8_t>((p ^ blind.a) & 1u);
|
||||
mine.e = static_cast<std::uint8_t>((q ^ blind.b) & 1u);
|
||||
gmw_and_msg theirs{};
|
||||
detail::exchange_pod(peer, peer_stream, mine, theirs);
|
||||
const std::uint8_t d_open =
|
||||
static_cast<std::uint8_t>(mine.d ^ theirs.d);
|
||||
const std::uint8_t e_open =
|
||||
static_cast<std::uint8_t>(mine.e ^ theirs.e);
|
||||
ot::bit_triple t{blind.a, blind.b, blind.c};
|
||||
return edabit::and_finish(t, d_open, e_open, party);
|
||||
}
|
||||
|
||||
/// @brief Mux: `b + sel·(a-b)` via one Beaver mul of `(sel, a-b)`.
|
||||
inline std::uint64_t mux_online(unsigned party, std::uint8_t sel,
|
||||
std::uint64_t a, std::uint64_t b, net::stream_array & peer,
|
||||
net::stream_array & dealer, std::size_t dealer_stream = 0,
|
||||
std::size_t peer_d_stream = 0, std::size_t peer_e_stream = 1,
|
||||
std::uint16_t limb = 8)
|
||||
{
|
||||
const auto mask = detail::limb_mask(limb);
|
||||
const auto diff = (a - b) & mask;
|
||||
const auto prod = beaver_mul_online(party, sel & 1u, diff, peer, dealer,
|
||||
dealer_stream, peer_d_stream, peer_e_stream, limb);
|
||||
return (b + prod) & mask;
|
||||
}
|
||||
|
||||
/// @brief Exact trunc: open `x-r` (r from `s` dabits), then GMW wrap via
|
||||
/// low-bit AND chain on the same shape as `a2b_online`'s carry.
|
||||
inline std::uint64_t trunc_exact_online(unsigned party, std::uint64_t x_share,
|
||||
unsigned n, unsigned s, net::stream_array & peer, net::stream_array & dealer,
|
||||
std::uint16_t limb = 8, std::size_t dabit_stream = 0,
|
||||
std::size_t and_blind_stream = 1, std::size_t open_stream = 0,
|
||||
std::size_t and_peer_stream = 1)
|
||||
{
|
||||
if (s >= n || n > 64)
|
||||
throw std::invalid_argument("trunc_exact_online");
|
||||
const auto mask = detail::limb_mask(limb);
|
||||
std::uint64_t r = 0;
|
||||
std::vector<std::uint8_t> r_bits(s);
|
||||
for (unsigned i = 0; i < s; ++i)
|
||||
{
|
||||
const auto d = detail::read_pod<dabit_view>(dealer, dabit_stream);
|
||||
r_bits[i] = static_cast<std::uint8_t>(d.bit & 1u);
|
||||
std::uint64_t ar = 0;
|
||||
const std::uint16_t use = d.limb != 0 ? d.limb : limb;
|
||||
std::memcpy(&ar, d.arith, use);
|
||||
r = (r + (ar << i)) & mask;
|
||||
}
|
||||
const auto delta =
|
||||
open_sum_online(party, (x_share - r) & mask, peer, open_stream, limb);
|
||||
// Wrap = carry out of the low `s` bits of delta + r (GMW chain).
|
||||
std::uint8_t carry = 0;
|
||||
for (unsigned i = 0; i < s; ++i)
|
||||
{
|
||||
const std::uint8_t di =
|
||||
static_cast<std::uint8_t>((delta >> i) & 1u);
|
||||
const std::uint8_t ri = r_bits[i];
|
||||
const std::uint8_t rc =
|
||||
gmw_and_online(party, ri, carry, peer, dealer, and_blind_stream,
|
||||
and_peer_stream);
|
||||
const std::uint8_t rd = static_cast<std::uint8_t>(di & ri);
|
||||
const std::uint8_t dc = static_cast<std::uint8_t>(di & carry);
|
||||
carry = static_cast<std::uint8_t>((rd ^ dc ^ rc) & 1u);
|
||||
}
|
||||
const auto d_wrap = detail::read_pod<dabit_view>(dealer, dabit_stream);
|
||||
const std::uint8_t mine_mask =
|
||||
static_cast<std::uint8_t>(carry ^ (d_wrap.bit & 1u));
|
||||
std::uint8_t peer_mask = 0;
|
||||
detail::exchange_pod(peer, open_stream, mine_mask, peer_mask);
|
||||
const std::uint8_t mask_open =
|
||||
static_cast<std::uint8_t>(mine_mask ^ peer_mask);
|
||||
ot::dabit<std::uint64_t> dr{};
|
||||
dr.bit = d_wrap.bit;
|
||||
const std::uint16_t use_wrap =
|
||||
d_wrap.limb != 0 ? d_wrap.limb : limb;
|
||||
std::memcpy(&dr.arith, d_wrap.arith, use_wrap);
|
||||
const auto wrap = edabit::b2a_party_bit<std::uint64_t>(
|
||||
carry, dr, mask_open, party, 0);
|
||||
const std::uint64_t high_m = ((n - s) >= 64)
|
||||
? ~std::uint64_t{0}
|
||||
: ((std::uint64_t{1} << (n - s)) - 1u);
|
||||
std::uint64_t out = wrap & mask;
|
||||
if (party == 0)
|
||||
out = (out + ((delta >> s) & high_m)) & mask;
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief A2B: `width` dabits on `dealer[dabit_stream]` (sequential reads),
|
||||
/// `width` AND blinds on `dealer[and_blind_stream]`, open delta on
|
||||
/// `peer[open_stream]`, AND masks on `peer[and_peer_stream]` (sequential).
|
||||
inline std::uint64_t a2b_online(unsigned party, std::uint64_t x_share,
|
||||
unsigned width, net::stream_array & peer, net::stream_array & dealer,
|
||||
std::uint16_t limb = 8, std::size_t dabit_stream = 0,
|
||||
std::size_t and_blind_stream = 1, std::size_t open_stream = 0,
|
||||
std::size_t and_peer_stream = 1)
|
||||
{
|
||||
if (width == 0 || width > 64)
|
||||
throw std::invalid_argument("a2b_online width");
|
||||
const auto mask = detail::limb_mask(limb);
|
||||
std::uint64_t r_arith = 0;
|
||||
std::vector<std::uint8_t> r_bits(width);
|
||||
for (unsigned i = 0; i < width; ++i)
|
||||
{
|
||||
const auto d = detail::read_pod<dabit_view>(dealer, dabit_stream);
|
||||
r_bits[i] = static_cast<std::uint8_t>(d.bit & 1u);
|
||||
std::uint64_t ar = 0;
|
||||
const std::uint16_t use = d.limb != 0 ? d.limb : limb;
|
||||
std::memcpy(&ar, d.arith, use);
|
||||
r_arith = (r_arith + (ar << i)) & mask;
|
||||
}
|
||||
const auto delta = open_sum_online(party, (x_share - r_arith) & mask, peer,
|
||||
open_stream, limb);
|
||||
std::uint8_t carry = 0;
|
||||
std::uint64_t out = 0;
|
||||
for (unsigned i = 0; i < width; ++i)
|
||||
{
|
||||
const auto blind =
|
||||
detail::read_pod<gmw_and_blind>(dealer, and_blind_stream);
|
||||
ot::bit_triple t{blind.a, blind.b, blind.c};
|
||||
const std::uint8_t di =
|
||||
static_cast<std::uint8_t>((delta >> i) & 1u);
|
||||
const std::uint8_t ri = r_bits[i];
|
||||
gmw_and_msg mine{};
|
||||
mine.d = static_cast<std::uint8_t>((ri ^ t.a) & 1u);
|
||||
mine.e = static_cast<std::uint8_t>((carry ^ t.b) & 1u);
|
||||
gmw_and_msg peer_msg{};
|
||||
detail::exchange_pod(peer, and_peer_stream, mine, peer_msg);
|
||||
const std::uint8_t d =
|
||||
static_cast<std::uint8_t>(mine.d ^ peer_msg.d);
|
||||
const std::uint8_t e =
|
||||
static_cast<std::uint8_t>(mine.e ^ peer_msg.e);
|
||||
const std::uint8_t rc = edabit::and_finish(t, d, e, party);
|
||||
const std::uint8_t bi = static_cast<std::uint8_t>(
|
||||
((party == 0 ? (ri ^ di ^ carry) : (ri ^ carry))) & 1u);
|
||||
const std::uint8_t rd = static_cast<std::uint8_t>(di & ri);
|
||||
const std::uint8_t dc = static_cast<std::uint8_t>(di & carry);
|
||||
carry = static_cast<std::uint8_t>((rd ^ dc ^ rc) & 1u);
|
||||
out |= (static_cast<std::uint64_t>(bi) << i);
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief Unsigned GT via two A2B passes plus MSB-first compare ANDs.
|
||||
inline std::uint8_t gt_online(unsigned party, std::uint64_t x, std::uint64_t y,
|
||||
unsigned width, net::stream_array & peer, net::stream_array & dealer,
|
||||
std::uint16_t limb = 8)
|
||||
{
|
||||
// Dealer layout: streams 0..1 dabits+ANDs for x, 2..3 for y, 4 for compare.
|
||||
const auto bx = a2b_online(party, x, width, peer, dealer, limb, 0, 1, 0, 1);
|
||||
const auto by = a2b_online(party, y, width, peer, dealer, limb, 2, 3, 2, 3);
|
||||
std::uint8_t gt = 0;
|
||||
std::uint8_t eq = static_cast<std::uint8_t>(party == 0 ? 1 : 0);
|
||||
for (unsigned k = width; k-- > 0; )
|
||||
{
|
||||
const std::uint8_t xk =
|
||||
static_cast<std::uint8_t>((bx >> k) & 1u);
|
||||
const std::uint8_t yk =
|
||||
static_cast<std::uint8_t>((by >> k) & 1u);
|
||||
const std::uint8_t xny = gmw_and_online(party, xk,
|
||||
static_cast<std::uint8_t>(yk ^ (party == 0 ? 1u : 0u)), peer,
|
||||
dealer, 4, 4);
|
||||
const std::uint8_t both = gmw_and_online(party, eq, xny, peer, dealer,
|
||||
4, 4);
|
||||
gt = static_cast<std::uint8_t>(gt ^ both);
|
||||
const std::uint8_t same = static_cast<std::uint8_t>(
|
||||
(xk ^ yk ^ (party == 0 ? 1u : 0u)) & 1u);
|
||||
eq = gmw_and_online(party, eq, same, peer, dealer, 4, 4);
|
||||
}
|
||||
return gt;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// RoundSink adapter: one peer stream index per circuit exchange round
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// @brief Drive `mpc::party::run` with one `stream_array` index per slot.
|
||||
class stream_array_circuit_sink : public net::RoundSink
|
||||
{
|
||||
public:
|
||||
stream_array_circuit_sink(net::stream_array & peer,
|
||||
std::vector<std::size_t> slot_bytes)
|
||||
: peer_(&peer), slot_bytes_(std::move(slot_bytes))
|
||||
{
|
||||
windows_.reserve(slot_bytes_.size());
|
||||
for (std::size_t sb : slot_bytes_)
|
||||
windows_.emplace_back(1, sb);
|
||||
if (peer_->size() < slot_bytes_.size())
|
||||
throw std::invalid_argument(
|
||||
"stream_array_circuit_sink: peer too few streams");
|
||||
}
|
||||
|
||||
std::size_t count() const noexcept override { return 1; }
|
||||
std::size_t rounds() const noexcept override { return slot_bytes_.size(); }
|
||||
|
||||
std::size_t slot_bytes(std::uint16_t round) const override
|
||||
{
|
||||
if (round >= slot_bytes_.size())
|
||||
throw std::out_of_range("stream_array_circuit_sink round");
|
||||
return slot_bytes_[round];
|
||||
}
|
||||
|
||||
void submit(std::uint16_t round, std::size_t index,
|
||||
const std::uint8_t * bytes, std::size_t n) override
|
||||
{
|
||||
if (index != 0)
|
||||
throw std::out_of_range("stream_array_circuit_sink index");
|
||||
window(round).submit(0, bytes, n);
|
||||
}
|
||||
|
||||
bool peer_ready(std::uint16_t round, std::size_t index) const override
|
||||
{
|
||||
if (index != 0)
|
||||
return false;
|
||||
return window(round).peer_ready(0);
|
||||
}
|
||||
|
||||
void read_peer(std::uint16_t round, std::size_t index, std::uint8_t * out,
|
||||
std::size_t n) const override
|
||||
{
|
||||
if (index != 0)
|
||||
throw std::out_of_range("stream_array_circuit_sink read");
|
||||
window(round).read_peer(0, out, n);
|
||||
}
|
||||
|
||||
void flush() override
|
||||
{
|
||||
for (std::uint16_t r = 0; r < slot_bytes_.size(); ++r)
|
||||
{
|
||||
std::size_t begin = 0;
|
||||
std::size_t n = 0;
|
||||
window(r).pending_out(begin, n);
|
||||
if (n != 0)
|
||||
flush_round(r);
|
||||
}
|
||||
}
|
||||
|
||||
void flush_round(std::uint16_t round) override
|
||||
{
|
||||
if (round >= slot_bytes_.size())
|
||||
throw std::out_of_range("stream_array_circuit_sink flush_round");
|
||||
auto & w = window(round);
|
||||
std::size_t begin = 0;
|
||||
std::size_t nslots = 0;
|
||||
const auto * pend = w.pending_out(begin, nslots);
|
||||
if (nslots == 0)
|
||||
return;
|
||||
const std::size_t nbyte = nslots * slot_bytes_[round];
|
||||
peer_->write(round, pend, nbyte);
|
||||
peer_->flush(round);
|
||||
w.mark_flushed(nslots);
|
||||
std::vector<std::uint8_t> peer_buf(nbyte);
|
||||
peer_->read(round, peer_buf.data(), nbyte);
|
||||
w.accept_peer_at(0, peer_buf.data(), nbyte);
|
||||
}
|
||||
|
||||
void poll() override {}
|
||||
|
||||
private:
|
||||
net::round_window & window(std::uint16_t round)
|
||||
{
|
||||
return windows_.at(round);
|
||||
}
|
||||
|
||||
const net::round_window & window(std::uint16_t round) const
|
||||
{
|
||||
return windows_.at(round);
|
||||
}
|
||||
|
||||
net::stream_array * peer_ = nullptr;
|
||||
std::vector<std::size_t> slot_bytes_;
|
||||
std::vector<net::round_window> windows_;
|
||||
};
|
||||
|
||||
/// @brief Run a compiled circuit using stream indices for each online slot.
|
||||
inline void run_circuit_stream_array(const mpc::circuit & circ, mpc::party & me,
|
||||
net::stream_array & peer, net::stream_array & dealer,
|
||||
std::size_t prep_bytes)
|
||||
{
|
||||
stream_array_circuit_sink sink(peer, circ.slot_bytes());
|
||||
me.run(sink, dealer, prep_bytes);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Gilboa OT pads on stream arrays
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
inline void write_ot_pack_wire(net::stream_array & dealer, std::size_t stream,
|
||||
const ot::pack::wire & w)
|
||||
{
|
||||
const std::uint64_t nb = static_cast<std::uint64_t>(w.b2a.size());
|
||||
const std::uint64_t nt = static_cast<std::uint64_t>(w.bits.size());
|
||||
const std::uint64_t nr =
|
||||
w.bit_ring_elem != 0
|
||||
? static_cast<std::uint64_t>(w.bit_ring.size() / w.bit_ring_elem)
|
||||
: 0u;
|
||||
const std::uint64_t hdr[4] = {static_cast<std::uint64_t>(w.me), nb, nt, nr};
|
||||
dealer.write(stream, hdr, sizeof(hdr));
|
||||
if (nb != 0)
|
||||
dealer.write(stream, w.b2a.data(), nb * sizeof(ot::b2a_slot));
|
||||
if (nt != 0)
|
||||
dealer.write(stream, w.bits.data(), nt * sizeof(ot::bit_triple));
|
||||
if (nr != 0)
|
||||
dealer.write(stream, w.bit_ring.data(), w.bit_ring.size());
|
||||
dealer.flush(stream);
|
||||
}
|
||||
|
||||
inline ot::pack read_ot_pack_wire(net::stream_array & dealer, std::size_t stream)
|
||||
{
|
||||
std::uint64_t hdr[4]{};
|
||||
dealer.read(stream, hdr, sizeof(hdr));
|
||||
ot::pack::wire w;
|
||||
w.me = static_cast<int>(hdr[0]);
|
||||
const std::size_t nb = static_cast<std::size_t>(hdr[1]);
|
||||
const std::size_t nt = static_cast<std::size_t>(hdr[2]);
|
||||
const std::size_t nr = static_cast<std::size_t>(hdr[3]);
|
||||
w.b2a.resize(nb);
|
||||
w.bits.resize(nt);
|
||||
if (nb != 0)
|
||||
dealer.read(stream, w.b2a.data(), nb * sizeof(ot::b2a_slot));
|
||||
if (nt != 0)
|
||||
dealer.read(stream, w.bits.data(), nt * sizeof(ot::bit_triple));
|
||||
if (nr != 0)
|
||||
{
|
||||
w.bit_ring_elem = sizeof(ot::bit_ring_triple<std::uint64_t>);
|
||||
w.bit_ring.resize(nr * w.bit_ring_elem);
|
||||
dealer.read(stream, w.bit_ring.data(), w.bit_ring.size());
|
||||
}
|
||||
return ot::pack::from_wire(std::move(w));
|
||||
}
|
||||
|
||||
/// @brief Sample correlated Gilboa/OT pads onto each party's dealer stream.
|
||||
inline void deal_gilboa_mul_tape(net::stream_array & dealer0,
|
||||
net::stream_array & dealer1, std::size_t bits)
|
||||
{
|
||||
auto pr = ot::sample_dealer_pair(bits, 0, bits);
|
||||
write_ot_pack_wire(dealer0, 0, pr.first.export_wire());
|
||||
write_ot_pack_wire(dealer1, 0, pr.second.export_wire());
|
||||
}
|
||||
|
||||
/// @brief Online Gilboa mul: OT pads on `dealer`, `d`/`e` vectors on `peer`.
|
||||
inline std::uint64_t gilboa_mul_online(unsigned party, std::uint64_t x,
|
||||
std::uint64_t y, net::stream_array & peer, net::stream_array & dealer,
|
||||
std::size_t bits = 16, std::size_t peer_stream = 0,
|
||||
std::size_t dealer_stream = 0)
|
||||
{
|
||||
auto pack = read_ot_pack_wire(dealer, dealer_stream);
|
||||
auto r = gilboa::mul_from_ot_begin(pack, x, y, party, static_cast<unsigned>(bits));
|
||||
std::vector<std::uint64_t> peer_d(r.d_share.size());
|
||||
std::vector<std::uint64_t> peer_e(r.e_share.size());
|
||||
for (std::size_t i = 0; i < r.d_share.size(); ++i)
|
||||
{
|
||||
detail::exchange_pod(peer, peer_stream, r.d_share[i], peer_d[i]);
|
||||
detail::exchange_pod(peer, peer_stream, r.e_share[i], peer_e[i]);
|
||||
}
|
||||
return gilboa::mul_from_ot_finish(r, peer_d, peer_e);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// RSS ring refresh on a 3-party stream clique
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// @brief Send `mine` to the next party; receive predecessor's payload.
|
||||
/// @details Writes `(own, from_prev)` into `own_out` / `next_out` (RSS pair).
|
||||
inline void rss_refresh_ring_online(net::memory_stream_clique & clique,
|
||||
unsigned me, const std::uint8_t * mine, std::size_t n,
|
||||
std::uint8_t * own_out, std::uint8_t * next_out, std::size_t stream = 0)
|
||||
{
|
||||
if (clique.parties != 3 || me > 2)
|
||||
throw std::invalid_argument("rss_refresh_ring_online");
|
||||
const unsigned next = static_cast<unsigned>((me + 1) % 3);
|
||||
const unsigned prev = static_cast<unsigned>((me + 2) % 3);
|
||||
clique.end(me, next).write(stream, mine, n);
|
||||
clique.end(me, next).flush(stream);
|
||||
clique.end(me, prev).read(stream, next_out, n);
|
||||
std::memcpy(own_out, mine, n);
|
||||
}
|
||||
|
||||
} // namespace factory
|
||||
} // namespace dpf
|
||||
|
||||
#endif
|
||||
107
include/dpf/factory_tapes.hpp
Normal file
107
include/dpf/factory_tapes.hpp
Normal file
|
|
@ -0,0 +1,107 @@
|
|||
/// @file dpf/factory_tapes.hpp
|
||||
/// @brief Dealer stream layouts for `factory_gadgets` online tapes (A2B, trunc, GT).
|
||||
#ifndef LIBDPF_INCLUDE_DPF_FACTORY_TAPES_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_FACTORY_TAPES_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <stdexcept>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
|
||||
#include "dpf/protocol_factory.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace factory
|
||||
{
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
inline std::pair<gmw_and_blind, gmw_and_blind> deal_gmw_and_blind_pair()
|
||||
{
|
||||
auto t = deal_bit_triple();
|
||||
gmw_and_blind v0{t.first.a, t.first.b, t.first.c};
|
||||
gmw_and_blind v1{t.second.a, t.second.b, t.second.c};
|
||||
return {v0, v1};
|
||||
}
|
||||
|
||||
template <typename F>
|
||||
inline void deal_tape_pod(net::stream_array & party0, net::stream_array & party1,
|
||||
std::size_t stream, std::size_t count, F && f)
|
||||
{
|
||||
if (stream >= party0.size() || stream >= party1.size())
|
||||
throw std::invalid_argument("deal_tape_pod: stream index");
|
||||
for (std::size_t j = 0; j < count; ++j)
|
||||
{
|
||||
auto views = f();
|
||||
using V0 = std::decay_t<decltype(views.first)>;
|
||||
using V1 = std::decay_t<decltype(views.second)>;
|
||||
static_assert(std::is_trivially_copyable_v<V0>
|
||||
&& std::is_trivially_copyable_v<V1>,
|
||||
"deal_tape_pod views must be trivially copyable");
|
||||
if (sizeof(V0) != sizeof(V1))
|
||||
throw std::logic_error("deal_tape_pod: unequal view sizes");
|
||||
detail::write_pod(party0, stream, views.first);
|
||||
detail::write_pod(party1, stream, views.second);
|
||||
}
|
||||
party0.flush(stream);
|
||||
party1.flush(stream);
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/// @brief Sequential dabits then AND blinds for one `a2b_online` pass.
|
||||
inline void deal_a2b_tape(unsigned width, std::uint16_t limb,
|
||||
net::stream_array & party0, net::stream_array & party1,
|
||||
std::size_t dabit_stream = 0, std::size_t and_blind_stream = 1)
|
||||
{
|
||||
if (width == 0 || width > 64)
|
||||
throw std::invalid_argument("deal_a2b_tape width");
|
||||
const std::size_t need =
|
||||
1 + (dabit_stream > and_blind_stream ? dabit_stream : and_blind_stream);
|
||||
if (party0.size() < need || party1.size() < need)
|
||||
throw std::invalid_argument("deal_a2b_tape: not enough streams");
|
||||
detail::deal_tape_pod(party0, party1, dabit_stream, width,
|
||||
[limb] { return deal_dabit(limb); });
|
||||
detail::deal_tape_pod(party0, party1, and_blind_stream, width,
|
||||
[] { return detail::deal_gmw_and_blind_pair(); });
|
||||
}
|
||||
|
||||
/// @brief `s + 1` dabits (low mask + B2A wrap) and `s` AND blinds for trunc.
|
||||
inline void deal_trunc_tape(unsigned n, unsigned s, std::uint16_t limb,
|
||||
net::stream_array & party0, net::stream_array & party1,
|
||||
std::size_t dabit_stream = 0, std::size_t and_blind_stream = 1)
|
||||
{
|
||||
if (s >= n || n > 64)
|
||||
throw std::invalid_argument("deal_trunc_tape");
|
||||
const std::size_t need =
|
||||
1 + (dabit_stream > and_blind_stream ? dabit_stream : and_blind_stream);
|
||||
if (party0.size() < need || party1.size() < need)
|
||||
throw std::invalid_argument("deal_trunc_tape: not enough streams");
|
||||
detail::deal_tape_pod(party0, party1, dabit_stream, s + 1u,
|
||||
[limb] { return deal_dabit(limb); });
|
||||
detail::deal_tape_pod(party0, party1, and_blind_stream, s,
|
||||
[] { return detail::deal_gmw_and_blind_pair(); });
|
||||
}
|
||||
|
||||
/// @brief Dealer layout for `gt_online`: streams 0–1 (x A2B), 2–3 (y), 4 compare.
|
||||
inline void deal_gt_tape(unsigned width, std::uint16_t limb,
|
||||
net::stream_array & party0, net::stream_array & party1)
|
||||
{
|
||||
if (width == 0 || width > 64)
|
||||
throw std::invalid_argument("deal_gt_tape width");
|
||||
constexpr std::size_t k_gt_streams = 5;
|
||||
if (party0.size() < k_gt_streams || party1.size() < k_gt_streams)
|
||||
throw std::invalid_argument("deal_gt_tape: need 5 dealer streams");
|
||||
deal_a2b_tape(width, limb, party0, party1, 0, 1);
|
||||
deal_a2b_tape(width, limb, party0, party1, 2, 3);
|
||||
detail::deal_tape_pod(party0, party1, 4, 3u * width,
|
||||
[] { return detail::deal_gmw_and_blind_pair(); });
|
||||
}
|
||||
|
||||
} // namespace factory
|
||||
} // namespace dpf
|
||||
|
||||
#endif
|
||||
492
include/dpf/field128.hpp
Normal file
492
include/dpf/field128.hpp
Normal file
|
|
@ -0,0 +1,492 @@
|
|||
/// @file dpf/field128.hpp
|
||||
/// @brief Prime field of libprio's `Field128`, as a DPF output.
|
||||
/// @details The modulus is 340282366920938462946865773367900766209. Pass
|
||||
/// `dpf::field128{n}` as a `make_dpf` payload. Leaf addition and
|
||||
/// scaling are the field operations. A raw PRG block is reduced into
|
||||
/// the field on the first leaf operation. This is a point-function
|
||||
/// output, not a comparison payload.
|
||||
/// @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_FIELD128_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_FIELD128_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <iomanip>
|
||||
#include <ostream>
|
||||
#include <type_traits>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
#include "simde/simde/x86/avx2.h"
|
||||
|
||||
#include "dpf/leaf_arithmetic.hpp"
|
||||
#include "dpf/random.hpp"
|
||||
#include "dpf/utils.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
/// @brief Element of the 128-bit libprio field.
|
||||
class field128
|
||||
{
|
||||
public:
|
||||
/// @brief Low limb of the modulus `2^128 - 0x1bffffffffffffffff`.
|
||||
static constexpr std::uint64_t mod_lo = 0x0000000000000001ull;
|
||||
/// @brief High limb of the modulus.
|
||||
static constexpr std::uint64_t mod_hi = 0xffffffffffffffe4ull;
|
||||
static constexpr bool dpf_point_group = true;
|
||||
|
||||
/// @brief The zero element.
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr field128() noexcept = default;
|
||||
|
||||
/// @brief Reduce `v` into the field. A negative value is negated in the field.
|
||||
/// @tparam T integral type, at most 128 bits
|
||||
/// @param v the integer to reduce
|
||||
template <typename T, typename = std::enable_if_t<std::is_integral_v<T>>>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr field128(T v) noexcept
|
||||
{
|
||||
assign_integer(v);
|
||||
}
|
||||
|
||||
/// @brief Reduce a PRG block into the field. Uses up to 16 bytes.
|
||||
/// @param bytes the PRG output
|
||||
/// @param n the number of bytes available
|
||||
/// @return the field element
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static field128 from_seed(const void * bytes, std::size_t n) noexcept
|
||||
{
|
||||
unsigned char buf[16]{};
|
||||
if (n > sizeof(buf))
|
||||
n = sizeof(buf);
|
||||
std::memcpy(buf, bytes, n);
|
||||
std::uint64_t w[2]{};
|
||||
std::memcpy(w, buf, sizeof(w));
|
||||
return from_lane(w[0], w[1]);
|
||||
}
|
||||
|
||||
/// @brief Reduce an unreduced 128-bit lane `(hi << 64) + lo`.
|
||||
/// @param lo the low limb, not necessarily reduced
|
||||
/// @param hi the high limb, not necessarily reduced
|
||||
/// @return the field element
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static constexpr field128 from_lane(std::uint64_t lo, std::uint64_t hi) noexcept
|
||||
{
|
||||
const std::uint64_t words[4] = {lo, hi, 0, 0};
|
||||
return reduce_words(words);
|
||||
}
|
||||
|
||||
/// @brief Canonical representative in `[0, mod)`.
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static constexpr field128 canonicalize(field128 a) noexcept
|
||||
{
|
||||
return from_lane(a.lo_, a.hi_);
|
||||
}
|
||||
|
||||
/// @brief Low 64 bits of the reduced representative.
|
||||
/// @return the low limb
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_PURE
|
||||
constexpr std::uint64_t lo() const noexcept { return lo_; }
|
||||
|
||||
/// @brief High 64 bits of the reduced representative.
|
||||
/// @return the high limb
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_PURE
|
||||
constexpr std::uint64_t hi() const noexcept { return hi_; }
|
||||
|
||||
/// @brief Field addition.
|
||||
/// @param a left addend
|
||||
/// @param b right addend
|
||||
/// @return `a + b` in the field
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
friend constexpr field128 operator+(field128 a, field128 b) noexcept
|
||||
{
|
||||
const unsigned __int128 low =
|
||||
static_cast<unsigned __int128>(a.lo_) + b.lo_;
|
||||
const unsigned __int128 high =
|
||||
static_cast<unsigned __int128>(a.hi_) + b.hi_ + (low >> 64);
|
||||
const std::uint64_t words[4] = {
|
||||
static_cast<std::uint64_t>(low),
|
||||
static_cast<std::uint64_t>(high),
|
||||
static_cast<std::uint64_t>(high >> 64),
|
||||
0};
|
||||
return reduce_words(words);
|
||||
}
|
||||
|
||||
/// @brief Field subtraction.
|
||||
/// @param a minuend
|
||||
/// @param b subtrahend
|
||||
/// @return `a - b` in the field
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
friend constexpr field128 operator-(field128 a, field128 b) noexcept
|
||||
{
|
||||
a = canonicalize(a);
|
||||
b = canonicalize(b);
|
||||
if (!less(a.hi_, a.lo_, b.hi_, b.lo_))
|
||||
{
|
||||
std::uint64_t lo = 0, hi = 0;
|
||||
sub_words(a.lo_, a.hi_, b.lo_, b.hi_, lo, hi);
|
||||
return from_reduced(lo, hi);
|
||||
}
|
||||
std::uint64_t lo = 0, hi = 0;
|
||||
sub_words(mod_lo, mod_hi, b.lo_, b.hi_, lo, hi);
|
||||
return from_reduced(lo, hi) + a;
|
||||
}
|
||||
|
||||
/// @brief Field negation.
|
||||
/// @param a the element to negate
|
||||
/// @return `-a`, with `-0 = 0`
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
friend constexpr field128 operator-(field128 a) noexcept
|
||||
{
|
||||
a = canonicalize(a);
|
||||
if (a.lo_ == 0 && a.hi_ == 0)
|
||||
return a;
|
||||
std::uint64_t lo = 0, hi = 0;
|
||||
sub_words(mod_lo, mod_hi, a.lo_, a.hi_, lo, hi);
|
||||
return from_reduced(lo, hi);
|
||||
}
|
||||
|
||||
/// @brief Field multiplication.
|
||||
/// @param a left factor
|
||||
/// @param b right factor
|
||||
/// @return `a * b` in the field
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
friend constexpr field128 operator*(field128 a, field128 b) noexcept
|
||||
{
|
||||
std::uint64_t words[4]{};
|
||||
mul128(a.lo_, a.hi_, b.lo_, b.hi_, words);
|
||||
return reduce_words(words);
|
||||
}
|
||||
|
||||
/// @brief Field equality.
|
||||
/// @param a left element
|
||||
/// @param b right element
|
||||
/// @return `true` when the reduced values match
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
friend constexpr bool operator==(field128 a, field128 b) noexcept
|
||||
{
|
||||
a = canonicalize(a);
|
||||
b = canonicalize(b);
|
||||
return a.lo_ == b.lo_ && a.hi_ == b.hi_;
|
||||
}
|
||||
|
||||
/// @brief Field inequality.
|
||||
/// @param a left element
|
||||
/// @param b right element
|
||||
/// @return `true` when the reduced values differ
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
friend constexpr bool operator!=(field128 a, field128 b) noexcept
|
||||
{
|
||||
return !(a == b);
|
||||
}
|
||||
|
||||
/// @brief Write the reduced representative in hexadecimal.
|
||||
/// @param os the output stream
|
||||
/// @param a the element to write
|
||||
/// @return `os`
|
||||
friend std::ostream & operator<<(std::ostream & os, field128 a)
|
||||
{
|
||||
const auto flags = os.flags();
|
||||
os << "0x" << std::hex << a.hi_ << std::setfill('0') << std::setw(16) << a.lo_;
|
||||
os.flags(flags);
|
||||
return os;
|
||||
}
|
||||
|
||||
private:
|
||||
std::uint64_t lo_{};
|
||||
std::uint64_t hi_{};
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static constexpr field128 from_reduced(std::uint64_t lo, std::uint64_t hi) noexcept
|
||||
{
|
||||
field128 out;
|
||||
out.lo_ = lo;
|
||||
out.hi_ = hi;
|
||||
return out;
|
||||
}
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static constexpr bool less(std::uint64_t hi, std::uint64_t lo,
|
||||
std::uint64_t ohi, std::uint64_t olo) noexcept
|
||||
{
|
||||
return hi < ohi || (hi == ohi && lo < olo);
|
||||
}
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static constexpr bool ge_mod(std::uint64_t hi, std::uint64_t lo) noexcept
|
||||
{
|
||||
return hi > mod_hi || (hi == mod_hi && lo >= mod_lo);
|
||||
}
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static constexpr void sub_words(std::uint64_t al, std::uint64_t ah,
|
||||
std::uint64_t bl, std::uint64_t bh,
|
||||
std::uint64_t &ol, std::uint64_t &oh) noexcept
|
||||
{
|
||||
const unsigned borrow = al < bl ? 1u : 0u;
|
||||
ol = static_cast<std::uint64_t>(al - bl);
|
||||
oh = static_cast<std::uint64_t>(ah - bh - borrow);
|
||||
}
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static constexpr void add_c(std::uint64_t &lo, std::uint64_t &hi,
|
||||
std::uint64_t &top) noexcept
|
||||
{
|
||||
const unsigned __int128 low =
|
||||
static_cast<unsigned __int128>(lo) + 0xffffffffffffffffull;
|
||||
lo = static_cast<std::uint64_t>(low);
|
||||
const unsigned __int128 high =
|
||||
static_cast<unsigned __int128>(hi) + 0x1bull + static_cast<std::uint64_t>(low >> 64);
|
||||
hi = static_cast<std::uint64_t>(high);
|
||||
top = static_cast<std::uint64_t>(high >> 64);
|
||||
}
|
||||
|
||||
/// @brief 128×128 → 256 multiply, little-endian limbs.
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static constexpr void mul128(std::uint64_t a0, std::uint64_t a1,
|
||||
std::uint64_t b0, std::uint64_t b1, std::uint64_t out[4]) noexcept
|
||||
{
|
||||
const unsigned __int128 p00 = static_cast<unsigned __int128>(a0) * b0;
|
||||
const unsigned __int128 p01 = static_cast<unsigned __int128>(a0) * b1;
|
||||
const unsigned __int128 p10 = static_cast<unsigned __int128>(a1) * b0;
|
||||
const unsigned __int128 p11 = static_cast<unsigned __int128>(a1) * b1;
|
||||
out[0] = static_cast<std::uint64_t>(p00);
|
||||
const unsigned __int128 mid = (p00 >> 64)
|
||||
+ static_cast<std::uint64_t>(p01) + static_cast<std::uint64_t>(p10);
|
||||
out[1] = static_cast<std::uint64_t>(mid);
|
||||
const unsigned __int128 high = (mid >> 64) + (p01 >> 64) + (p10 >> 64)
|
||||
+ static_cast<std::uint64_t>(p11);
|
||||
out[2] = static_cast<std::uint64_t>(high);
|
||||
out[3] = static_cast<std::uint64_t>((high >> 64) + (p11 >> 64));
|
||||
}
|
||||
|
||||
/// @brief Reduce a little-endian 256-bit integer. `2^128 ≡ C`.
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static constexpr field128 reduce_words(const std::uint64_t z[4]) noexcept
|
||||
{
|
||||
std::uint64_t lo = 0, hi = 0, top = 0;
|
||||
for (int bit = 255; bit >= 0; --bit)
|
||||
{
|
||||
top = hi >> 63;
|
||||
hi = (hi << 1) | (lo >> 63);
|
||||
lo <<= 1;
|
||||
const unsigned limb = static_cast<unsigned>(bit >> 6);
|
||||
const unsigned off = static_cast<unsigned>(bit & 63);
|
||||
if ((z[limb] >> off) & 1u)
|
||||
{
|
||||
const unsigned __int128 low =
|
||||
static_cast<unsigned __int128>(lo) + 1;
|
||||
lo = static_cast<std::uint64_t>(low);
|
||||
const unsigned __int128 high =
|
||||
static_cast<unsigned __int128>(hi) + static_cast<std::uint64_t>(low >> 64);
|
||||
hi = static_cast<std::uint64_t>(high);
|
||||
top += static_cast<std::uint64_t>(high >> 64);
|
||||
}
|
||||
while (top)
|
||||
add_c(lo, hi, top);
|
||||
if (ge_mod(hi, lo))
|
||||
sub_words(lo, hi, mod_lo, mod_hi, lo, hi);
|
||||
}
|
||||
return from_reduced(lo, hi);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr void assign_integer(T v) noexcept
|
||||
{
|
||||
bool neg = false;
|
||||
unsigned __int128 mag = 0;
|
||||
if constexpr (std::is_signed_v<T>)
|
||||
{
|
||||
if (v < 0)
|
||||
{
|
||||
neg = true;
|
||||
using U = std::make_unsigned_t<T>;
|
||||
mag = static_cast<U>(0) - static_cast<U>(v);
|
||||
}
|
||||
else
|
||||
{
|
||||
mag = static_cast<std::make_unsigned_t<T>>(v);
|
||||
}
|
||||
}
|
||||
else
|
||||
{
|
||||
mag = static_cast<unsigned __int128>(v);
|
||||
}
|
||||
const std::uint64_t words[4] = {
|
||||
static_cast<std::uint64_t>(mag),
|
||||
static_cast<std::uint64_t>(mag >> 64),
|
||||
0, 0};
|
||||
*this = reduce_words(words);
|
||||
if (neg)
|
||||
*this = -*this;
|
||||
}
|
||||
};
|
||||
|
||||
namespace utils
|
||||
{
|
||||
|
||||
template <>
|
||||
struct bitlength_of<field128>
|
||||
: std::integral_constant<std::size_t, 128>
|
||||
{ };
|
||||
|
||||
template <>
|
||||
struct has_characteristic_two<field128> : std::false_type
|
||||
{ };
|
||||
|
||||
} // namespace utils
|
||||
|
||||
namespace leaf_arithmetic
|
||||
{
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
field128 load_field128(const void *lane) noexcept
|
||||
{
|
||||
std::uint64_t w[2]{};
|
||||
std::memcpy(w, lane, sizeof(w));
|
||||
return field128::from_lane(w[0], w[1]);
|
||||
}
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void store_field128(void *lane, field128 v) noexcept
|
||||
{
|
||||
const std::uint64_t w[2] = {v.lo(), v.hi()};
|
||||
std::memcpy(lane, w, sizeof(w));
|
||||
}
|
||||
|
||||
template <std::size_t Lanes>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void field128_lanes(const void *a, const void *b, void *out,
|
||||
field128 (*op)(field128, field128)) noexcept
|
||||
{
|
||||
auto *aa = static_cast<const unsigned char *>(a);
|
||||
auto *bb = static_cast<const unsigned char *>(b);
|
||||
auto *cc = static_cast<unsigned char *>(out);
|
||||
for (std::size_t i = 0; i < Lanes; ++i)
|
||||
{
|
||||
const field128 y = op(load_field128(aa + i * 16), load_field128(bb + i * 16));
|
||||
store_field128(cc + i * 16, y);
|
||||
}
|
||||
}
|
||||
|
||||
template <std::size_t Lanes>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void field128_scale(const void *a, field128 b, void *out) noexcept
|
||||
{
|
||||
auto *aa = static_cast<const unsigned char *>(a);
|
||||
auto *cc = static_cast<unsigned char *>(out);
|
||||
for (std::size_t i = 0; i < Lanes; ++i)
|
||||
store_field128(cc + i * 16, load_field128(aa + i * 16) * b);
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
HEDLEY_PRAGMA(GCC diagnostic push)
|
||||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
||||
|
||||
template <>
|
||||
struct add_t<field128, simde__m128i>
|
||||
{
|
||||
auto operator()(const simde__m128i &a, const simde__m128i &b) const
|
||||
{
|
||||
simde__m128i out;
|
||||
detail::field128_lanes<1>(&a, &b, &out, [](field128 x, field128 y) {
|
||||
return x + y;
|
||||
});
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct subtract_t<field128, simde__m128i>
|
||||
{
|
||||
auto operator()(const simde__m128i &a, const simde__m128i &b) const
|
||||
{
|
||||
simde__m128i out;
|
||||
detail::field128_lanes<1>(&a, &b, &out, [](field128 x, field128 y) {
|
||||
return x - y;
|
||||
});
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct multiply_t<field128, simde__m128i>
|
||||
{
|
||||
auto operator()(const simde__m128i &a, field128 b) const
|
||||
{
|
||||
simde__m128i out;
|
||||
detail::field128_scale<1>(&a, b, &out);
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct add_t<field128, simde__m256i>
|
||||
{
|
||||
auto operator()(const simde__m256i &a, const simde__m256i &b) const
|
||||
{
|
||||
simde__m256i out;
|
||||
detail::field128_lanes<2>(&a, &b, &out, [](field128 x, field128 y) {
|
||||
return x + y;
|
||||
});
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct subtract_t<field128, simde__m256i>
|
||||
{
|
||||
auto operator()(const simde__m256i &a, const simde__m256i &b) const
|
||||
{
|
||||
simde__m256i out;
|
||||
detail::field128_lanes<2>(&a, &b, &out, [](field128 x, field128 y) {
|
||||
return x - y;
|
||||
});
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct multiply_t<field128, simde__m256i>
|
||||
{
|
||||
auto operator()(const simde__m256i &a, field128 b) const
|
||||
{
|
||||
simde__m256i out;
|
||||
detail::field128_scale<2>(&a, b, &out);
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||||
|
||||
} // namespace leaf_arithmetic
|
||||
|
||||
/// @brief Sample a uniform field element by rejection.
|
||||
/// @return an element of the field
|
||||
template <>
|
||||
HEDLEY_NO_THROW
|
||||
inline auto uniform_sample<field128>() noexcept
|
||||
{
|
||||
for (;;)
|
||||
{
|
||||
const auto lo = uniform_sample<std::uint64_t>();
|
||||
const auto hi = uniform_sample<std::uint64_t>();
|
||||
if (hi < field128::mod_hi || (hi == field128::mod_hi && lo < field128::mod_lo))
|
||||
return field128::from_lane(lo, hi);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_FIELD128_HPP__
|
||||
362
include/dpf/field64.hpp
Normal file
362
include/dpf/field64.hpp
Normal file
|
|
@ -0,0 +1,362 @@
|
|||
/// @file dpf/field64.hpp
|
||||
/// @brief Prime field GF(2^64 − 2^32 + 1), as a DPF output.
|
||||
/// @details The modulus is libprio's `Field64` (the Goldilocks prime). Pass
|
||||
/// `dpf::field64{n}` as a `make_dpf` payload. Leaf addition and
|
||||
/// scaling are the field operations. A raw PRG block is reduced into
|
||||
/// the field on the first leaf operation. This is a point-function
|
||||
/// output, not a comparison payload.
|
||||
/// @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_FIELD64_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_FIELD64_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <ostream>
|
||||
#include <type_traits>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
#include "simde/simde/x86/avx2.h"
|
||||
|
||||
#include "dpf/leaf_arithmetic.hpp"
|
||||
#include "dpf/random.hpp"
|
||||
#include "dpf/utils.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
/// @brief Element of GF(2^64 − 2^32 + 1).
|
||||
class field64
|
||||
{
|
||||
public:
|
||||
/// @brief Modulus \f$2^{64}-2^{32}+1\f$.
|
||||
static constexpr std::uint64_t mod = 0xffffffff00000001ull;
|
||||
using integral_type = std::uint64_t;
|
||||
static constexpr bool dpf_point_group = true;
|
||||
|
||||
/// @brief Reduce `v` into the field. A negative value is negated in the field.
|
||||
/// @tparam T integral type
|
||||
/// @param v the integer to reduce
|
||||
template <typename T, typename = std::enable_if_t<std::is_integral_v<T>>>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr field64(T v) noexcept
|
||||
{
|
||||
val = from_integer(v);
|
||||
}
|
||||
|
||||
/// @brief The zero element.
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr field64() noexcept = default;
|
||||
|
||||
/// @brief Reduce a PRG block into the field. Uses up to 16 bytes.
|
||||
/// @param bytes the PRG output
|
||||
/// @param n the number of bytes available
|
||||
/// @return the field element
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static field64 from_seed(const void * bytes, std::size_t n) noexcept
|
||||
{
|
||||
unsigned char buf[16]{};
|
||||
if (n > sizeof(buf))
|
||||
n = sizeof(buf);
|
||||
std::memcpy(buf, bytes, n);
|
||||
std::uint64_t w[2]{};
|
||||
std::memcpy(w, buf, sizeof(w));
|
||||
const unsigned __int128 wide = static_cast<unsigned __int128>(w[0])
|
||||
| (static_cast<unsigned __int128>(w[1]) << 64);
|
||||
return from_reduced(reduce_u128(wide));
|
||||
}
|
||||
|
||||
/// @brief Reduced representative in `[0, mod)`.
|
||||
/// @return the stored field element
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_PURE
|
||||
constexpr integral_type raw() const noexcept { return reduce(val); }
|
||||
|
||||
/// @brief Same value as `raw()`.
|
||||
/// @return the stored field element
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_PURE
|
||||
explicit constexpr operator integral_type() const noexcept { return raw(); }
|
||||
|
||||
/// @brief Fold a 128-bit integer into the field.
|
||||
/// @param x the integer to reduce
|
||||
/// @return `x` modulo `mod`
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static constexpr integral_type reduce_u128(unsigned __int128 x) noexcept
|
||||
{
|
||||
// 2^64 ≡ 2^32 − 1, so each step replaces the high half.
|
||||
while (x >> 64)
|
||||
{
|
||||
const auto lo = static_cast<integral_type>(x);
|
||||
const auto hi = static_cast<integral_type>(x >> 64);
|
||||
x = static_cast<unsigned __int128>(lo)
|
||||
+ (static_cast<unsigned __int128>(hi) << 32) - hi;
|
||||
}
|
||||
auto lo = static_cast<integral_type>(x);
|
||||
if (lo >= mod)
|
||||
lo -= mod;
|
||||
return lo;
|
||||
}
|
||||
|
||||
/// @brief Reduce a 64-bit word. `mod` itself fits in the word.
|
||||
/// @param x the integer to reduce
|
||||
/// @return `x` modulo `mod`
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static constexpr integral_type reduce(integral_type x) noexcept
|
||||
{
|
||||
return x >= mod ? static_cast<integral_type>(x - mod) : x;
|
||||
}
|
||||
|
||||
/// @brief Field addition.
|
||||
/// @param a left addend
|
||||
/// @param b right addend
|
||||
/// @return `a + b` in the field
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
friend constexpr field64 operator+(field64 a, field64 b) noexcept
|
||||
{
|
||||
return from_reduced(reduce_u128(
|
||||
static_cast<unsigned __int128>(reduce(a.val)) + reduce(b.val)));
|
||||
}
|
||||
|
||||
/// @brief Field subtraction.
|
||||
/// @param a minuend
|
||||
/// @param b subtrahend
|
||||
/// @return `a - b` in the field
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
friend constexpr field64 operator-(field64 a, field64 b) noexcept
|
||||
{
|
||||
a.val = reduce(a.val);
|
||||
b.val = reduce(b.val);
|
||||
if (a.val >= b.val)
|
||||
return from_reduced(static_cast<integral_type>(a.val - b.val));
|
||||
return from_reduced(static_cast<integral_type>(mod - (b.val - a.val)));
|
||||
}
|
||||
|
||||
/// @brief Field negation.
|
||||
/// @param a the element to negate
|
||||
/// @return `-a`, with `-0 = 0`
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
friend constexpr field64 operator-(field64 a) noexcept
|
||||
{
|
||||
a.val = reduce(a.val);
|
||||
return from_reduced(a.val == 0 ? 0 : static_cast<integral_type>(mod - a.val));
|
||||
}
|
||||
|
||||
/// @brief Field multiplication.
|
||||
/// @param a left factor
|
||||
/// @param b right factor
|
||||
/// @return `a * b` in the field
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
friend constexpr field64 operator*(field64 a, field64 b) noexcept
|
||||
{
|
||||
return from_reduced(reduce_u128(
|
||||
static_cast<unsigned __int128>(reduce(a.val)) * reduce(b.val)));
|
||||
}
|
||||
|
||||
/// @brief Field equality.
|
||||
/// @param a left element
|
||||
/// @param b right element
|
||||
/// @return `true` when the reduced values match
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
friend constexpr bool operator==(field64 a, field64 b) noexcept
|
||||
{
|
||||
return reduce(a.val) == reduce(b.val);
|
||||
}
|
||||
|
||||
/// @brief Field inequality.
|
||||
/// @param a left element
|
||||
/// @param b right element
|
||||
/// @return `true` when the reduced values differ
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
friend constexpr bool operator!=(field64 a, field64 b) noexcept
|
||||
{
|
||||
return reduce(a.val) != reduce(b.val);
|
||||
}
|
||||
|
||||
/// @brief Write the reduced representative in decimal.
|
||||
/// @param os the output stream
|
||||
/// @param a the element to write
|
||||
/// @return `os`
|
||||
friend std::ostream & operator<<(std::ostream & os, field64 a)
|
||||
{
|
||||
return os << a.raw();
|
||||
}
|
||||
|
||||
private:
|
||||
integral_type val{};
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static constexpr field64 from_reduced(integral_type v) noexcept
|
||||
{
|
||||
field64 out;
|
||||
out.val = v;
|
||||
return out;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static constexpr integral_type from_integer(T v) noexcept
|
||||
{
|
||||
if constexpr (std::is_signed_v<T>)
|
||||
{
|
||||
if (v < 0)
|
||||
{
|
||||
using U = std::make_unsigned_t<T>;
|
||||
const auto mag = static_cast<U>(0) - static_cast<U>(v);
|
||||
return reduce_u128(static_cast<unsigned __int128>(mag)) == 0
|
||||
? 0
|
||||
: static_cast<integral_type>(
|
||||
mod - reduce_u128(static_cast<unsigned __int128>(mag)));
|
||||
}
|
||||
}
|
||||
return reduce_u128(static_cast<unsigned __int128>(v));
|
||||
}
|
||||
};
|
||||
|
||||
namespace utils
|
||||
{
|
||||
|
||||
template <>
|
||||
struct bitlength_of<field64>
|
||||
: std::integral_constant<std::size_t, 64>
|
||||
{ };
|
||||
|
||||
template <>
|
||||
struct has_characteristic_two<field64> : std::false_type
|
||||
{ };
|
||||
|
||||
} // namespace utils
|
||||
|
||||
namespace leaf_arithmetic
|
||||
{
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
template <std::size_t Lanes>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void field64_lanes(const void *a, const void *b, void *out,
|
||||
field64 (*op)(field64, field64)) noexcept
|
||||
{
|
||||
std::uint64_t aa[Lanes], bb[Lanes], cc[Lanes];
|
||||
std::memcpy(aa, a, sizeof(aa));
|
||||
std::memcpy(bb, b, sizeof(bb));
|
||||
for (std::size_t i = 0; i < Lanes; ++i)
|
||||
cc[i] = op(field64{aa[i]}, field64{bb[i]}).raw();
|
||||
std::memcpy(out, cc, sizeof(cc));
|
||||
}
|
||||
|
||||
template <std::size_t Lanes>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void field64_scale(const void *a, field64 b, void *out) noexcept
|
||||
{
|
||||
std::uint64_t aa[Lanes], cc[Lanes];
|
||||
std::memcpy(aa, a, sizeof(aa));
|
||||
for (std::size_t i = 0; i < Lanes; ++i)
|
||||
cc[i] = (field64{aa[i]} * b).raw();
|
||||
std::memcpy(out, cc, sizeof(cc));
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
HEDLEY_PRAGMA(GCC diagnostic push)
|
||||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
||||
|
||||
template <>
|
||||
struct add_t<field64, simde__m128i>
|
||||
{
|
||||
auto operator()(const simde__m128i &a, const simde__m128i &b) const
|
||||
{
|
||||
simde__m128i out;
|
||||
detail::field64_lanes<2>(&a, &b, &out, [](field64 x, field64 y) {
|
||||
return x + y;
|
||||
});
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct subtract_t<field64, simde__m128i>
|
||||
{
|
||||
auto operator()(const simde__m128i &a, const simde__m128i &b) const
|
||||
{
|
||||
simde__m128i out;
|
||||
detail::field64_lanes<2>(&a, &b, &out, [](field64 x, field64 y) {
|
||||
return x - y;
|
||||
});
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct multiply_t<field64, simde__m128i>
|
||||
{
|
||||
auto operator()(const simde__m128i &a, field64 b) const
|
||||
{
|
||||
simde__m128i out;
|
||||
detail::field64_scale<2>(&a, b, &out);
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct add_t<field64, simde__m256i>
|
||||
{
|
||||
auto operator()(const simde__m256i &a, const simde__m256i &b) const
|
||||
{
|
||||
simde__m256i out;
|
||||
detail::field64_lanes<4>(&a, &b, &out, [](field64 x, field64 y) {
|
||||
return x + y;
|
||||
});
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct subtract_t<field64, simde__m256i>
|
||||
{
|
||||
auto operator()(const simde__m256i &a, const simde__m256i &b) const
|
||||
{
|
||||
simde__m256i out;
|
||||
detail::field64_lanes<4>(&a, &b, &out, [](field64 x, field64 y) {
|
||||
return x - y;
|
||||
});
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct multiply_t<field64, simde__m256i>
|
||||
{
|
||||
auto operator()(const simde__m256i &a, field64 b) const
|
||||
{
|
||||
simde__m256i out;
|
||||
detail::field64_scale<4>(&a, b, &out);
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||||
|
||||
} // namespace leaf_arithmetic
|
||||
|
||||
/// @brief Sample a uniform field element by rejection.
|
||||
/// @return an element of the field
|
||||
template <>
|
||||
HEDLEY_NO_THROW
|
||||
inline auto uniform_sample<field64>() noexcept
|
||||
{
|
||||
for (;;)
|
||||
{
|
||||
const auto x = uniform_sample<std::uint64_t>();
|
||||
if (x < field64::mod)
|
||||
return field64{x};
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_FIELD64_HPP__
|
||||
87
include/dpf/fixed_share.hpp
Normal file
87
include/dpf/fixed_share.hpp
Normal file
|
|
@ -0,0 +1,87 @@
|
|||
/// @file dpf/fixed_share.hpp
|
||||
/// @brief Fixed-point arithmetic shares: `fixed<IntBits, FracBits>`.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_FIXED_SHARE_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_FIXED_SHARE_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <stdexcept>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/share_vec.hpp"
|
||||
#include "dpf/trunc.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
template <unsigned IntBits, unsigned FracBits>
|
||||
struct fixed
|
||||
{
|
||||
static constexpr unsigned int_bits = IntBits;
|
||||
static constexpr unsigned frac_bits = FracBits;
|
||||
static constexpr unsigned width = IntBits + FracBits;
|
||||
static_assert(width > 0 && width <= 128, "fixed width 1..128");
|
||||
|
||||
using ring = std::conditional_t<(width > 64), unsigned __int128, std::uint64_t>;
|
||||
|
||||
ring raw{};
|
||||
|
||||
fixed() = default;
|
||||
explicit fixed(ring v) : raw(v) {}
|
||||
|
||||
static fixed from_integer(ring i)
|
||||
{
|
||||
return fixed{static_cast<ring>(i << FracBits)};
|
||||
}
|
||||
|
||||
static fixed mul_clear(fixed a, fixed b)
|
||||
{
|
||||
const ring prod = static_cast<ring>(a.raw * b.raw);
|
||||
return fixed{static_cast<ring>(prod >> FracBits)};
|
||||
}
|
||||
|
||||
friend fixed operator+(fixed a, fixed b)
|
||||
{
|
||||
return fixed{static_cast<ring>(a.raw + b.raw)};
|
||||
}
|
||||
|
||||
friend fixed operator-(fixed a, fixed b)
|
||||
{
|
||||
return fixed{static_cast<ring>(a.raw - b.raw)};
|
||||
}
|
||||
};
|
||||
|
||||
template <unsigned IntBits, unsigned FracBits>
|
||||
using fixed_vec = share_vec<typename fixed<IntBits, FracBits>::ring>;
|
||||
|
||||
/// @brief Fixed-point product via mul_trunc by FracBits (Beaver, both parties).
|
||||
template <unsigned IntBits, unsigned FracBits>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::pair<fixed_vec<IntBits, FracBits>, fixed_vec<IntBits, FracBits>>
|
||||
fixed_mul_share(const fixed_vec<IntBits, FracBits> & x0,
|
||||
const fixed_vec<IntBits, FracBits> & x1,
|
||||
const fixed_vec<IntBits, FracBits> & y0,
|
||||
const fixed_vec<IntBits, FracBits> & y1)
|
||||
{
|
||||
if (x0.size() != x1.size() || x0.size() != y0.size()
|
||||
|| y0.size() != y1.size())
|
||||
throw std::invalid_argument("fixed_mul_share size");
|
||||
fixed_vec<IntBits, FracBits> z0(x0.size(), protocol::domain::a, 0);
|
||||
fixed_vec<IntBits, FracBits> z1(x0.size(), protocol::domain::a, 1);
|
||||
for (std::size_t i = 0; i < x0.size(); ++i)
|
||||
{
|
||||
// Widen so IntBits+FracBits products do not wrap before the shift.
|
||||
auto mt = trunc::mul_exact_trunc(x0[i], x1[i], y0[i], y1[i], FracBits);
|
||||
z0[i] = mt.z0;
|
||||
z1[i] = mt.z1;
|
||||
}
|
||||
return {std::move(z0), std::move(z1)};
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_FIXED_SHARE_HPP__
|
||||
260
include/dpf/flute.hpp
Normal file
260
include/dpf/flute.hpp
Normal file
|
|
@ -0,0 +1,260 @@
|
|||
/// @file dpf/flute.hpp
|
||||
/// @brief Lookup tables as a multi-fan-in inner product.
|
||||
/// @details A public table `f : {0,1}^δ → {0,1}^σ` is the OR of the input
|
||||
/// rows whose output bit is 1. Those terms are disjoint, so the OR
|
||||
/// is an XOR, and the XOR of ANDs is an inner product of one vector
|
||||
/// per input bit (and the public output column). Complements flip the
|
||||
/// already-opened masked bit and leave the mask λ alone, so setup
|
||||
/// builds subset products of the δ input masks once, not once per row.
|
||||
///
|
||||
/// Online, each party XORs the public terms into a share of `v` and
|
||||
/// the parties exchange that share. Two parties send two bits per
|
||||
/// output bit, independent of δ. The masked output `m_z` is then
|
||||
/// local, and the semantic bit is `m_z XOR λ_z`.
|
||||
///
|
||||
/// The two-party exchange is ΠLUT. The three-party function is the
|
||||
/// same algebra on three XOR shares of each mask: the paper's
|
||||
/// construction is the inner product, and the share count is the
|
||||
/// underlying MPC. This is not a new share domain.
|
||||
/// @note Andreas Brüggemann, Robin Hundt, Thomas Schneider, Ajith Suresh, and
|
||||
/// Hossein Yalame, "FLUTE: Fast and Secure Lookup Table Evaluations,"
|
||||
/// IEEE S&P 2023 (ePrint 2023/499). The masked bits are the ABY2.0 wire
|
||||
/// (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_FLUTE_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_FLUTE_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <stdexcept>
|
||||
#include <vector>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/random.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace flute
|
||||
{
|
||||
|
||||
inline constexpr unsigned k_max_delta = 8;
|
||||
|
||||
/// @brief One evaluation. `opened[w] = masked[w] XOR` the mask shares.
|
||||
struct result
|
||||
{
|
||||
std::vector<std::uint8_t> opened;
|
||||
std::vector<std::uint8_t> masked;
|
||||
/// @brief XOR shares of `λ_z`, one vector per party. Size 2 or 3.
|
||||
std::vector<std::vector<std::uint8_t>> mask;
|
||||
/// @brief Bits exchanged online: one bit per party per output bit.
|
||||
std::size_t online_bits = 0;
|
||||
};
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
inline std::uint8_t bit_of(std::uint32_t row, unsigned i)
|
||||
{
|
||||
return static_cast<std::uint8_t>((row >> i) & 1u);
|
||||
}
|
||||
|
||||
/// @brief Masked literal of input `i` on row `j`. Complement flips `m` only.
|
||||
inline std::uint8_t masked_lit(std::uint8_t m, std::uint8_t encoding)
|
||||
{
|
||||
return encoding ? m : static_cast<std::uint8_t>(m ^ 1u);
|
||||
}
|
||||
|
||||
inline std::uint8_t and_subset(std::uint32_t subset, std::uint32_t row,
|
||||
unsigned delta, const std::uint8_t * m)
|
||||
{
|
||||
std::uint8_t acc = 1;
|
||||
for (unsigned i = 0; i < delta; ++i)
|
||||
{
|
||||
if (((subset >> i) & 1u) == 0)
|
||||
continue;
|
||||
acc = static_cast<std::uint8_t>(
|
||||
acc & masked_lit(m[i], bit_of(row, i)));
|
||||
}
|
||||
return acc;
|
||||
}
|
||||
|
||||
inline std::uint8_t dot_column(std::uint32_t subset, unsigned delta,
|
||||
const std::uint8_t * m, const std::uint8_t * column)
|
||||
{
|
||||
const std::uint32_t rows = 1u << delta;
|
||||
std::uint8_t acc = 0;
|
||||
for (std::uint32_t j = 0; j < rows; ++j)
|
||||
{
|
||||
if (column[j] == 0)
|
||||
continue;
|
||||
acc = static_cast<std::uint8_t>(
|
||||
acc ^ and_subset(subset, j, delta, m));
|
||||
}
|
||||
return acc;
|
||||
}
|
||||
|
||||
inline std::uint8_t rand_bit()
|
||||
{
|
||||
return static_cast<std::uint8_t>(dpf::uniform_sample<std::uint8_t>() & 1u);
|
||||
}
|
||||
|
||||
struct setup
|
||||
{
|
||||
std::vector<std::uint8_t> m;
|
||||
std::vector<std::uint8_t> lam;
|
||||
/// @brief `share[party][subset] =` that party's XOR share of AND_{i in subset} λ_i.
|
||||
std::vector<std::vector<std::uint8_t>> share;
|
||||
std::vector<std::vector<std::uint8_t>> lamz;
|
||||
};
|
||||
|
||||
inline setup make_setup(unsigned delta, unsigned n_out, unsigned parties,
|
||||
const std::uint8_t * x)
|
||||
{
|
||||
if (parties != 2 && parties != 3)
|
||||
throw std::invalid_argument("flute: parties");
|
||||
setup s;
|
||||
s.m.resize(delta);
|
||||
s.lam.resize(delta);
|
||||
s.share.assign(parties, std::vector<std::uint8_t>(1u << delta, 0));
|
||||
s.lam.resize(delta);
|
||||
for (unsigned i = 0; i < delta; ++i)
|
||||
{
|
||||
if (x[i] > 1)
|
||||
throw std::invalid_argument("flute: input bit");
|
||||
std::uint8_t lam = 0;
|
||||
for (unsigned p = 0; p < parties; ++p)
|
||||
{
|
||||
const std::uint8_t sh = rand_bit();
|
||||
s.share[p][1u << i] = sh;
|
||||
lam = static_cast<std::uint8_t>(lam ^ sh);
|
||||
}
|
||||
s.lam[i] = lam;
|
||||
s.m[i] = static_cast<std::uint8_t>(x[i] ^ lam);
|
||||
}
|
||||
s.share[0][0] = 1;
|
||||
const std::uint32_t nsub = 1u << delta;
|
||||
for (std::uint32_t subset = 1; subset < nsub; ++subset)
|
||||
{
|
||||
if ((subset & (subset - 1u)) == 0)
|
||||
continue;
|
||||
for (unsigned p = 0; p < parties; ++p)
|
||||
{
|
||||
// Dealer shares the AND. Party 0 holds the product of the full
|
||||
// masks adjusted by the other parties' shares of this subset.
|
||||
s.share[p][subset] = 0;
|
||||
}
|
||||
std::uint8_t prod = 1;
|
||||
for (unsigned i = 0; i < delta; ++i)
|
||||
if (((subset >> i) & 1u) != 0)
|
||||
prod = static_cast<std::uint8_t>(prod & s.lam[i]);
|
||||
for (unsigned p = 1; p < parties; ++p)
|
||||
s.share[p][subset] = rand_bit();
|
||||
std::uint8_t rest = prod;
|
||||
for (unsigned p = 1; p < parties; ++p)
|
||||
rest = static_cast<std::uint8_t>(rest ^ s.share[p][subset]);
|
||||
s.share[0][subset] = rest;
|
||||
}
|
||||
s.lamz.assign(parties, std::vector<std::uint8_t>(n_out, 0));
|
||||
for (unsigned w = 0; w < n_out; ++w)
|
||||
for (unsigned p = 0; p < parties; ++p)
|
||||
s.lamz[p][w] = rand_bit();
|
||||
return s;
|
||||
}
|
||||
|
||||
inline std::uint8_t party_v(const setup & s, unsigned party, unsigned delta,
|
||||
std::uint32_t full, const std::uint8_t * column, std::uint8_t lamz_share)
|
||||
{
|
||||
std::uint8_t v = lamz_share;
|
||||
for (std::uint32_t subset = 0; subset < full; ++subset)
|
||||
{
|
||||
const std::uint32_t rest = full ^ subset;
|
||||
const std::uint8_t t = dot_column(subset, delta, s.m.data(), column);
|
||||
v = static_cast<std::uint8_t>(
|
||||
v ^ (t & s.share[party][rest]));
|
||||
}
|
||||
return v;
|
||||
}
|
||||
|
||||
inline result run(unsigned delta, unsigned n_out, unsigned parties,
|
||||
const std::uint8_t * columns, const std::uint8_t * x)
|
||||
{
|
||||
if (delta < 1 || delta > k_max_delta)
|
||||
throw std::invalid_argument("flute: delta");
|
||||
if (n_out == 0 || columns == nullptr || x == nullptr)
|
||||
throw std::invalid_argument("flute: table");
|
||||
const std::uint32_t rows = 1u << delta;
|
||||
const std::uint32_t full = rows - 1u;
|
||||
auto s = make_setup(delta, n_out, parties, x);
|
||||
result out;
|
||||
out.mask.assign(parties, std::vector<std::uint8_t>(n_out, 0));
|
||||
out.masked.resize(n_out);
|
||||
out.opened.resize(n_out);
|
||||
out.online_bits = static_cast<std::size_t>(parties) * n_out;
|
||||
for (unsigned w = 0; w < n_out; ++w)
|
||||
{
|
||||
const std::uint8_t * column = columns + static_cast<std::size_t>(w) * rows;
|
||||
std::uint8_t v = 0;
|
||||
for (unsigned p = 0; p < parties; ++p)
|
||||
{
|
||||
const std::uint8_t share = party_v(s, p, delta, full, column, s.lamz[p][w]);
|
||||
v = static_cast<std::uint8_t>(v ^ share);
|
||||
out.mask[p][w] = s.lamz[p][w];
|
||||
}
|
||||
const std::uint8_t t_full = dot_column(full, delta, s.m.data(), column);
|
||||
out.masked[w] = static_cast<std::uint8_t>(v ^ t_full);
|
||||
std::uint8_t lam = 0;
|
||||
for (unsigned p = 0; p < parties; ++p)
|
||||
lam = static_cast<std::uint8_t>(lam ^ out.mask[p][w]);
|
||||
out.opened[w] = static_cast<std::uint8_t>(out.masked[w] ^ lam);
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/// @brief Clear table. `x` is δ bits, least-significant bit first.
|
||||
/// `columns[w * 2^δ + row]` is output bit `w` on that row.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline std::vector<std::uint8_t> eval_plain(unsigned delta, unsigned n_out,
|
||||
const std::uint8_t * columns, const std::uint8_t * x)
|
||||
{
|
||||
if (delta < 1 || delta > k_max_delta)
|
||||
throw std::invalid_argument("flute: delta");
|
||||
std::uint32_t row = 0;
|
||||
for (unsigned i = 0; i < delta; ++i)
|
||||
{
|
||||
if (x[i] > 1)
|
||||
throw std::invalid_argument("flute: input bit");
|
||||
row |= static_cast<std::uint32_t>(x[i]) << i;
|
||||
}
|
||||
const std::uint32_t rows = 1u << delta;
|
||||
std::vector<std::uint8_t> out(n_out);
|
||||
for (unsigned w = 0; w < n_out; ++w)
|
||||
out[w] = columns[static_cast<std::size_t>(w) * rows + row] & 1u;
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief Two-party ΠLUT. Online cost is two bits per output bit.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline result eval_pair(unsigned delta, unsigned n_out,
|
||||
const std::uint8_t * columns, const std::uint8_t * x)
|
||||
{
|
||||
return detail::run(delta, n_out, 2, columns, x);
|
||||
}
|
||||
|
||||
/// @brief The same ΠLUT inner product on three XOR shares of each mask.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline result eval_trio(unsigned delta, unsigned n_out,
|
||||
const std::uint8_t * columns, const std::uint8_t * x)
|
||||
{
|
||||
return detail::run(delta, n_out, 3, columns, x);
|
||||
}
|
||||
|
||||
} // namespace flute
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_FLUTE_HPP__
|
||||
|
|
@ -11,17 +11,19 @@
|
|||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <ostream>
|
||||
#include <stdexcept>
|
||||
#include <type_traits>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/utils.hpp"
|
||||
#include "dpf/leaf_arithmetic.hpp"
|
||||
#include "dpf/random.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
/// @brief Modulus \(p = 2^{61}-1\). `p` itself reduces to 0.
|
||||
/// @brief Modulus \f$p = 2^{61}-1\f$. `p` itself reduces to 0.
|
||||
inline constexpr std::uint64_t fp61_mod = (std::uint64_t{1} << 61) - 1;
|
||||
|
||||
/// @brief Additive element of the field of order `2^61 - 1`.
|
||||
|
|
@ -66,7 +68,29 @@ class fp61
|
|||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_PURE
|
||||
constexpr integral_type raw() const noexcept { return val; }
|
||||
constexpr integral_type raw() const noexcept { return reduce(val); }
|
||||
|
||||
/// @brief Build a field element from PRG bytes (Mersenne reduction).
|
||||
/// @param bytes the PRG output
|
||||
/// @param n the number of bytes available
|
||||
/// @return the field element
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static fp61 from_seed(const void * bytes, std::size_t n) noexcept
|
||||
{
|
||||
unsigned char buf[16]{};
|
||||
if (n > sizeof(buf))
|
||||
n = sizeof(buf);
|
||||
std::memcpy(buf, bytes, n);
|
||||
std::uint64_t w[2]{};
|
||||
std::memcpy(w, buf, sizeof(w));
|
||||
using u128 = unsigned __int128;
|
||||
const u128 wide = static_cast<u128>(w[0])
|
||||
| (static_cast<u128>(w[1]) << 64);
|
||||
const auto lo = static_cast<integral_type>(wide) & fp61_mod;
|
||||
const auto mid = static_cast<integral_type>(wide >> 61) & fp61_mod;
|
||||
const auto hi = static_cast<integral_type>(wide >> 122);
|
||||
return fp61{lo + mid + hi};
|
||||
}
|
||||
|
||||
/// @brief Same value as `raw()`.
|
||||
/// @return the stored field element
|
||||
|
|
@ -98,7 +122,7 @@ class fp61
|
|||
HEDLEY_CONST
|
||||
friend constexpr fp61 operator+(fp61 a, fp61 b) noexcept
|
||||
{
|
||||
return fp61{a.val + b.val};
|
||||
return fp61{reduce(a.val) + reduce(b.val)};
|
||||
}
|
||||
|
||||
/// @brief Field subtraction.
|
||||
|
|
@ -110,7 +134,7 @@ class fp61
|
|||
HEDLEY_CONST
|
||||
friend constexpr fp61 operator-(fp61 a, fp61 b) noexcept
|
||||
{
|
||||
return fp61{a.val + fp61_mod - b.val};
|
||||
return fp61{reduce(a.val) + fp61_mod - reduce(b.val)};
|
||||
}
|
||||
|
||||
/// @brief Field negation.
|
||||
|
|
@ -121,7 +145,8 @@ class fp61
|
|||
HEDLEY_CONST
|
||||
friend constexpr fp61 operator-(fp61 a) noexcept
|
||||
{
|
||||
return fp61{a.val == 0 ? 0 : fp61_mod - a.val};
|
||||
const auto v = reduce(a.val);
|
||||
return fp61{v == 0 ? 0 : fp61_mod - v};
|
||||
}
|
||||
|
||||
/// @brief Field multiplication.
|
||||
|
|
@ -134,7 +159,7 @@ class fp61
|
|||
friend constexpr fp61 operator*(fp61 a, fp61 b) noexcept
|
||||
{
|
||||
using u128 = unsigned __int128;
|
||||
const u128 p = static_cast<u128>(a.val) * static_cast<u128>(b.val);
|
||||
const u128 p = static_cast<u128>(reduce(a.val)) * static_cast<u128>(reduce(b.val));
|
||||
const auto lo = static_cast<integral_type>(p) & fp61_mod;
|
||||
const auto mid = static_cast<integral_type>(p >> 61) & fp61_mod;
|
||||
const auto hi = static_cast<integral_type>(p >> 122);
|
||||
|
|
@ -150,7 +175,7 @@ class fp61
|
|||
HEDLEY_CONST
|
||||
friend constexpr bool operator==(fp61 a, fp61 b) noexcept
|
||||
{
|
||||
return a.val == b.val;
|
||||
return reduce(a.val) == reduce(b.val);
|
||||
}
|
||||
|
||||
/// @brief Field inequality.
|
||||
|
|
@ -162,7 +187,7 @@ class fp61
|
|||
HEDLEY_CONST
|
||||
friend constexpr bool operator!=(fp61 a, fp61 b) noexcept
|
||||
{
|
||||
return a.val != b.val;
|
||||
return reduce(a.val) != reduce(b.val);
|
||||
}
|
||||
|
||||
/// @brief Write the reduced representative in decimal.
|
||||
|
|
@ -195,6 +220,35 @@ struct has_characteristic_two<fp61> : std::false_type
|
|||
namespace leaf_arithmetic
|
||||
{
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
template <std::size_t Lanes>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void fp61_lanes(const void * a, const void * b, void * out,
|
||||
fp61 (*op)(fp61, fp61)) noexcept
|
||||
{
|
||||
std::uint64_t aa[Lanes], bb[Lanes], cc[Lanes];
|
||||
std::memcpy(aa, a, sizeof(aa));
|
||||
std::memcpy(bb, b, sizeof(bb));
|
||||
for (std::size_t i = 0; i < Lanes; ++i)
|
||||
cc[i] = op(fp61{aa[i]}, fp61{bb[i]}).raw();
|
||||
std::memcpy(out, cc, sizeof(cc));
|
||||
}
|
||||
|
||||
template <std::size_t Lanes>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void fp61_scale(const void * a, fp61 b, void * out) noexcept
|
||||
{
|
||||
std::uint64_t aa[Lanes], cc[Lanes];
|
||||
std::memcpy(aa, a, sizeof(aa));
|
||||
for (std::size_t i = 0; i < Lanes; ++i)
|
||||
cc[i] = (fp61{aa[i]} * b).raw();
|
||||
std::memcpy(out, cc, sizeof(cc));
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
HEDLEY_PRAGMA(GCC diagnostic push)
|
||||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
||||
template <>
|
||||
|
|
@ -202,7 +256,11 @@ struct add_t<fp61, simde__m128i>
|
|||
{
|
||||
auto operator()(const simde__m128i & a, const simde__m128i & b) const
|
||||
{
|
||||
return add_t<fp61::integral_type, simde__m128i>{}(a, b);
|
||||
simde__m128i out;
|
||||
detail::fp61_lanes<2>(&a, &b, &out, [](fp61 x, fp61 y) {
|
||||
return x + y;
|
||||
});
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
|
|
@ -211,7 +269,22 @@ struct subtract_t<fp61, simde__m128i>
|
|||
{
|
||||
auto operator()(const simde__m128i & a, const simde__m128i & b) const
|
||||
{
|
||||
return subtract_t<fp61::integral_type, simde__m128i>{}(a, b);
|
||||
simde__m128i out;
|
||||
detail::fp61_lanes<2>(&a, &b, &out, [](fp61 x, fp61 y) {
|
||||
return x - y;
|
||||
});
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct multiply_t<fp61, simde__m128i>
|
||||
{
|
||||
auto operator()(const simde__m128i & a, fp61 b) const
|
||||
{
|
||||
simde__m128i out;
|
||||
detail::fp61_scale<2>(&a, b, &out);
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
|
|
@ -220,7 +293,11 @@ struct add_t<fp61, simde__m256i>
|
|||
{
|
||||
auto operator()(const simde__m256i & a, const simde__m256i & b) const
|
||||
{
|
||||
return add_t<fp61::integral_type, simde__m256i>{}(a, b);
|
||||
simde__m256i out;
|
||||
detail::fp61_lanes<4>(&a, &b, &out, [](fp61 x, fp61 y) {
|
||||
return x + y;
|
||||
});
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
|
|
@ -229,13 +306,72 @@ struct subtract_t<fp61, simde__m256i>
|
|||
{
|
||||
auto operator()(const simde__m256i & a, const simde__m256i & b) const
|
||||
{
|
||||
return subtract_t<fp61::integral_type, simde__m256i>{}(a, b);
|
||||
simde__m256i out;
|
||||
detail::fp61_lanes<4>(&a, &b, &out, [](fp61 x, fp61 y) {
|
||||
return x - y;
|
||||
});
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct multiply_t<fp61, simde__m256i>
|
||||
{
|
||||
auto operator()(const simde__m256i & a, fp61 b) const
|
||||
{
|
||||
simde__m256i out;
|
||||
detail::fp61_scale<4>(&a, b, &out);
|
||||
return out;
|
||||
}
|
||||
};
|
||||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||||
|
||||
} // namespace leaf_arithmetic
|
||||
|
||||
/// @brief Sample a uniformly reduced field element by rejection.
|
||||
/// @return an element of the field
|
||||
template <>
|
||||
HEDLEY_NO_THROW
|
||||
inline auto uniform_sample<fp61>() noexcept
|
||||
{
|
||||
for (;;)
|
||||
{
|
||||
const auto v = uniform_sample<std::uint64_t>() & fp61_mod;
|
||||
if (v < fp61_mod)
|
||||
return fp61{v};
|
||||
}
|
||||
}
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
template <>
|
||||
struct shamir_field<fp61> : std::true_type
|
||||
{
|
||||
/// @brief `a^{-1}` by Fermat, `a^{p-2}`.
|
||||
/// @param a a non-zero field element
|
||||
/// @return `a^{-1}`
|
||||
/// @throws std::invalid_argument if `a` is zero
|
||||
static fp61 inv(fp61 a)
|
||||
{
|
||||
if (a.raw() == 0)
|
||||
throw std::invalid_argument("shamir: inverse of zero");
|
||||
fp61 base = a;
|
||||
fp61 out{1};
|
||||
auto e = fp61_mod - 2;
|
||||
while (e != 0)
|
||||
{
|
||||
if (e & 1u)
|
||||
out = out * base;
|
||||
base = base * base;
|
||||
e >>= 1;
|
||||
}
|
||||
return out;
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace detail
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_FP61_HPP__
|
||||
|
|
|
|||
|
|
@ -11,9 +11,10 @@
|
|||
/// cancels.
|
||||
///
|
||||
/// Default calls take XOR shares of the point. Tagged with
|
||||
/// `arith_input`, the point is the ring sum of the two shares; path
|
||||
/// bits are opened by a carry chain inside the local CW protocol so
|
||||
/// the words match `make_dpf(x0 + x1)` at the caller's query.
|
||||
/// `arith_input`, the point is the ring sum of the two shares. A
|
||||
/// beaver ripple-carry converts those shares to XOR shares of the
|
||||
/// sum bits before the walk, so the words match `make_dpf` on that
|
||||
/// sum. The sum is not opened.
|
||||
///
|
||||
/// `geneval_cmp` is the comparison-channel form. The value-correction
|
||||
/// word is a function of the secret path at every level, so the walk
|
||||
|
|
@ -21,6 +22,7 @@
|
|||
/// Doerner–Shelat comparison key. Prefix shares are
|
||||
/// `eval_point(cmp, ...)` at each endpoint. Piecewise-cubic evaluation
|
||||
/// on top of that is `grotto::geneval_offset_horner`.
|
||||
/// @note The per-level correction opening follows Jack Doerner and abhi shelat, CCS 2017 (ePrint 2017/827). They return a reusable key. This function opens a word only for nodes on the public query trie and, with a local pad tape, sends nothing.
|
||||
/// @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.
|
||||
|
|
@ -46,6 +48,7 @@
|
|||
#include "dpf/doerner_shelat.hpp"
|
||||
#include "dpf/eval_target.hpp"
|
||||
#include "dpf/leaf_node.hpp"
|
||||
#include "dpf/verifiable.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
|
@ -69,6 +72,9 @@ struct geneval_result
|
|||
std::size_t live_levels = 0;
|
||||
bool leaf_live = false;
|
||||
Leaf leaf{};
|
||||
/// @brief Party 0 / 1 VDPF tokens over the live eval trie (empty when unused).
|
||||
proof_token proof0{};
|
||||
proof_token proof1{};
|
||||
};
|
||||
|
||||
namespace detail
|
||||
|
|
@ -172,7 +178,8 @@ auto geneval_run(bool arith, bool arith_out, InputT x0, InputT x1,
|
|||
InputT x0c = x0;
|
||||
InputT x1c = x1;
|
||||
proto.encode_walk_shares(x0c, x1c, arith);
|
||||
const InputT alpha = utils::xor_input_shares(x0c, x1c);
|
||||
// Keep the secret path on share-bits. Do not form a clear alpha for leaf
|
||||
// placement, live levels, or correction seeds.
|
||||
|
||||
std::vector<InputT> flipped;
|
||||
flipped.reserve(queries.size());
|
||||
|
|
@ -191,8 +198,6 @@ auto geneval_run(bool arith, bool arith_out, InputT x0, InputT x1,
|
|||
if (unique_leaves.size() > (std::size_t{1} << 20))
|
||||
throw std::length_error("geneval trie is too large");
|
||||
|
||||
const uint64_t secret_leaf = geneval_leaf_id<dpf_type>(alpha);
|
||||
|
||||
constexpr auto to_int = utils::to_integral_type<InputT>{};
|
||||
|
||||
using tree = dpf::tree_traits<InteriorPRG>;
|
||||
|
|
@ -222,14 +227,17 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
std::memset(&result.leaf, 0, sizeof(result.leaf));
|
||||
result.correction_words.reserve(depth);
|
||||
result.correction_advice.reserve(depth);
|
||||
result.proof0 = detail::vdpf::zero_proof();
|
||||
result.proof1 = detail::vdpf::zero_proof();
|
||||
|
||||
auto mask = dpf_type::msb_mask;
|
||||
bool still_live = true;
|
||||
uint64_t secret_prefix = 0;
|
||||
for (std::size_t level = 0; level < depth; ++level, mask >>= 1)
|
||||
{
|
||||
const uint8_t bit0 = static_cast<uint8_t>(!!(to_int(mask) & to_int(x0c)));
|
||||
const uint8_t bit1 = static_cast<uint8_t>(!!(to_int(mask) & to_int(x1c)));
|
||||
const uint64_t parent_id = geneval_prefix(secret_leaf, depth, level);
|
||||
const uint64_t parent_id = secret_prefix;
|
||||
const bool is_last = tree::is_last_level(level, depth);
|
||||
|
||||
node L0 = simde_mm_setzero_si128();
|
||||
|
|
@ -303,6 +311,8 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
const node cw0 = tree::pack_cw(cw, advice, false, is_last);
|
||||
const node cw1 = tree::pack_cw(cw, advice, true, is_last);
|
||||
const std::size_t child_bits = level + 1;
|
||||
const uint8_t secret_bit = static_cast<uint8_t>((bit0 ^ bit1) & 1u);
|
||||
secret_prefix = (secret_prefix << 1) | secret_bit;
|
||||
std::vector<slot> next;
|
||||
next.reserve(exps.size() * 2);
|
||||
for (const exp & e : exps)
|
||||
|
|
@ -322,9 +332,37 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
dpf::xor_if_lo_bit(e.R1, cw1, e.s1)});
|
||||
}
|
||||
}
|
||||
|
||||
// Fold every live child into both parties' VDPF tokens.
|
||||
if (!next.empty())
|
||||
{
|
||||
cs_block cs{};
|
||||
bool have_cs = false;
|
||||
for (const slot & c : next)
|
||||
{
|
||||
if (c.id == secret_prefix)
|
||||
{
|
||||
// Prefix is the share-bit path accumulated above — not a
|
||||
// fresh xor_input_shares of the point for leaf placement.
|
||||
cs = detail::vdpf::make_cs(level, c.id, c.s0, c.s1);
|
||||
have_cs = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if (!have_cs)
|
||||
cs = detail::vdpf::make_cs(level, next[0].id, next[0].s0,
|
||||
next[0].s1);
|
||||
for (const slot & c : next)
|
||||
{
|
||||
detail::vdpf::fold_node(result.proof0, level, c.id, c.s0, cs);
|
||||
detail::vdpf::fold_node(result.proof1, level, c.id, c.s1, cs);
|
||||
}
|
||||
}
|
||||
|
||||
frontier = std::move(next);
|
||||
}
|
||||
|
||||
const uint64_t secret_leaf = secret_prefix;
|
||||
result.leaf_live = geneval_any_prefix(unique_leaves, depth, secret_leaf, depth);
|
||||
if (result.leaf_live)
|
||||
{
|
||||
|
|
@ -343,18 +381,21 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
{
|
||||
const uint8_t t0 = static_cast<uint8_t>(dpf::get_lo_bit(on->s0));
|
||||
const uint8_t t1 = static_cast<uint8_t>(dpf::get_lo_bit(on->s1));
|
||||
const std::size_t lane = static_cast<std::size_t>(to_int(alpha));
|
||||
result.leaf = proto.template open_arith_leaf<ExteriorPRG, 0, outputs_tuple>(
|
||||
dpf::unset_lo_2bits(on->s0), dpf::unset_lo_2bits(on->s1), t0, t1,
|
||||
y0, y1, std::size_t{0}, lane);
|
||||
y0, y1, std::size_t{0}, x0c, x1c);
|
||||
}
|
||||
else
|
||||
{
|
||||
const bool sign0 = dpf::get_lo_bit(on->s0);
|
||||
auto built = dpf::make_leaves<ExteriorPRG>(alpha,
|
||||
dpf::unset_lo_2bits(on->s0), dpf::unset_lo_2bits(on->s1), sign0,
|
||||
std::size_t{0}, y0);
|
||||
result.leaf = std::get<0>(built.first.first);
|
||||
// Mux / reconstruct only inside the leaf protocol hook.
|
||||
proto.open_leaf_group(x0c, x1c, [&](InputT sx0, InputT sx1) {
|
||||
const InputT x = utils::xor_input_shares(sx0, sx1);
|
||||
const bool sign0 = dpf::get_lo_bit(on->s0);
|
||||
auto built = dpf::make_leaves<ExteriorPRG>(x,
|
||||
dpf::unset_lo_2bits(on->s0), dpf::unset_lo_2bits(on->s1),
|
||||
sign0, std::size_t{0}, y0);
|
||||
result.leaf = std::get<0>(built.first.first);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -491,6 +532,10 @@ std::vector<InputT> geneval_inclusive(InputT from, InputT to)
|
|||
/// @param rng the Doerner–Shelat randomness tapes
|
||||
/// @param y the payload
|
||||
/// @return the opened shares and correction words
|
||||
/// \complexity O(n F) PRG expansions. n is `depth`. Each level expands every frontier node (two `expand` calls, one per share) and, while the secret path is live, one `prepare_level`. F is at most the number of distinct query leaves; the function rejects more than 2^20. Setup sorts the q query ids.
|
||||
/// \rounds No sockets. While the path is live, each level calls `prepare_level`, which samples one `ds_cw_pads` and two `ds_and_pads` and then `open_cw`.
|
||||
/// \communication none in this function.
|
||||
/// \preprocessing The `ds_randomness` tape: per live level, `ds_sample_cw` draws two 128-bit blocks and two bits, and each of the two AND pads is one `ds_sample_and`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
|
|
@ -512,6 +557,10 @@ auto geneval_point(InputT x0, InputT x1, InputT query,
|
|||
/// @param rng the Doerner–Shelat randomness tapes
|
||||
/// @param y the payload
|
||||
/// @return the opened shares and correction words
|
||||
/// \complexity O(n F) PRG expansions. n is `depth`. Each level expands every frontier node (two `expand` calls, one per share) and, while the secret path is live, one `prepare_level`. F is at most the number of distinct query leaves; the function rejects more than 2^20. Setup sorts the q query ids.
|
||||
/// \rounds No sockets. While the path is live, each level calls `prepare_level`, which samples one `ds_cw_pads` and two `ds_and_pads` and then `open_cw`.
|
||||
/// \communication none in this function.
|
||||
/// \preprocessing The `ds_randomness` tape: per live level, `ds_sample_cw` draws two 128-bit blocks and two bits, and each of the two AND pads is one `ds_sample_and`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
|
|
@ -534,6 +583,10 @@ auto geneval_point(arith_input_t, InputT x0, InputT x1, InputT query,
|
|||
/// @param y0 party 0's share of the payload
|
||||
/// @param y1 party 1's share of the payload
|
||||
/// @return the opened shares and correction words
|
||||
/// \complexity O(n F) PRG expansions. n is `depth`. Each level expands every frontier node (two `expand` calls, one per share) and, while the secret path is live, one `prepare_level`. F is at most the number of distinct query leaves; the function rejects more than 2^20. Setup sorts the q query ids.
|
||||
/// \rounds No sockets. While the path is live, each level calls `prepare_level`, which samples one `ds_cw_pads` and two `ds_and_pads` and then `open_cw`.
|
||||
/// \communication none in this function.
|
||||
/// \preprocessing The `ds_randomness` tape: per live level, `ds_sample_cw` draws two 128-bit blocks and two bits, and each of the two AND pads is one `ds_sample_and`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
|
|
@ -556,6 +609,10 @@ auto geneval_point(arith_output_t, InputT x0, InputT x1, InputT query,
|
|||
/// @param y0 party 0's share of the payload
|
||||
/// @param y1 party 1's share of the payload
|
||||
/// @return the opened shares and correction words
|
||||
/// \complexity O(n F) PRG expansions. n is `depth`. Each level expands every frontier node (two `expand` calls, one per share) and, while the secret path is live, one `prepare_level`. F is at most the number of distinct query leaves; the function rejects more than 2^20. Setup sorts the q query ids.
|
||||
/// \rounds No sockets. While the path is live, each level calls `prepare_level`, which samples one `ds_cw_pads` and two `ds_and_pads` and then `open_cw`.
|
||||
/// \communication none in this function.
|
||||
/// \preprocessing The `ds_randomness` tape: per live level, `ds_sample_cw` draws two 128-bit blocks and two bits, and each of the two AND pads is one `ds_sample_and`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
|
|
@ -586,6 +643,10 @@ auto geneval_point(arith_input_t, arith_output_t, InputT x0, InputT x1,
|
|||
/// @param rng the Doerner–Shelat randomness tapes
|
||||
/// @param y the `y`
|
||||
/// @return Geneval on the inclusive interval `[from, to]`
|
||||
/// \complexity O(n F) PRG expansions. n is `depth`. Each level expands every frontier node (two `expand` calls, one per share) and, while the secret path is live, one `prepare_level`. F is at most the number of distinct query leaves; the function rejects more than 2^20. Setup sorts the q query ids.
|
||||
/// \rounds No sockets. While the path is live, each level calls `prepare_level`, which samples one `ds_cw_pads` and two `ds_and_pads` and then `open_cw`.
|
||||
/// \communication none in this function.
|
||||
/// \preprocessing The `ds_randomness` tape: per live level, `ds_sample_cw` draws two 128-bit blocks and two bits, and each of the two AND pads is one `ds_sample_and`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
|
|
@ -600,6 +661,10 @@ auto geneval_interval(InputT x0, InputT x1, InputT from, InputT to,
|
|||
detail::geneval_inclusive(from, to), rng.root, rng.pad, y);
|
||||
}
|
||||
|
||||
/// \complexity O(n F) PRG expansions. n is `depth`. Each level expands every frontier node (two `expand` calls, one per share) and, while the secret path is live, one `prepare_level`. F is at most the number of distinct query leaves; the function rejects more than 2^20. Setup sorts the q query ids.
|
||||
/// \rounds No sockets. While the path is live, each level calls `prepare_level`, which samples one `ds_cw_pads` and two `ds_and_pads` and then `open_cw`.
|
||||
/// \communication none in this function.
|
||||
/// \preprocessing The `ds_randomness` tape: per live level, `ds_sample_cw` draws two 128-bit blocks and two bits, and each of the two AND pads is one `ds_sample_and`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
|
|
@ -614,7 +679,10 @@ auto geneval_interval(arith_input_t, InputT x0, InputT x1, InputT from,
|
|||
detail::geneval_inclusive(from, to), rng.root, rng.pad, y);
|
||||
}
|
||||
|
||||
/// @brief Geneval on the whole domain. Refuses a domain above 2^20 inputs.
|
||||
/// @brief Geneval on the whole domain.
|
||||
/// @details Materializes the query list. Domains wider than 20 bits refuse so
|
||||
/// a caller does not allocate a `2^n` vector by accident. Prefer the
|
||||
/// buffer overloads when writing into a pre-sized output scratch.
|
||||
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
|
||||
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
|
||||
/// @tparam InputT input domain type
|
||||
|
|
@ -626,6 +694,10 @@ auto geneval_interval(arith_input_t, InputT x0, InputT x1, InputT from,
|
|||
/// @param rng the Doerner–Shelat randomness tapes
|
||||
/// @param y the `y`
|
||||
/// @return Geneval on the whole domain
|
||||
/// \complexity O(n F) PRG expansions. n is `depth`. Each level expands every frontier node (two `expand` calls, one per share) and, while the secret path is live, one `prepare_level`. F is at most the number of distinct query leaves; the function rejects more than 2^20. Setup sorts the q query ids.
|
||||
/// \rounds No sockets. While the path is live, each level calls `prepare_level`, which samples one `ds_cw_pads` and two `ds_and_pads` and then `open_cw`.
|
||||
/// \communication none in this function.
|
||||
/// \preprocessing The `ds_randomness` tape: per live level, `ds_sample_cw` draws two 128-bit blocks and two bits, and each of the two AND pads is one `ds_sample_and`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
|
|
@ -640,6 +712,10 @@ auto geneval_full(InputT x0, InputT x1,
|
|||
detail::geneval_full_domain<InputT>(), rng.root, rng.pad, y);
|
||||
}
|
||||
|
||||
/// \complexity O(n F) PRG expansions. n is `depth`. Each level expands every frontier node (two `expand` calls, one per share) and, while the secret path is live, one `prepare_level`. F is at most the number of distinct query leaves; the function rejects more than 2^20. Setup sorts the q query ids.
|
||||
/// \rounds No sockets. While the path is live, each level calls `prepare_level`, which samples one `ds_cw_pads` and two `ds_and_pads` and then `open_cw`.
|
||||
/// \communication none in this function.
|
||||
/// \preprocessing The `ds_randomness` tape: per live level, `ds_sample_cw` draws two 128-bit blocks and two bits, and each of the two AND pads is one `ds_sample_and`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
|
|
@ -669,6 +745,10 @@ auto geneval_full(arith_input_t, InputT x0, InputT x1,
|
|||
/// @param rng the Doerner–Shelat randomness tapes
|
||||
/// @param y the `y`
|
||||
/// @return Geneval on a public sequence, in the order given
|
||||
/// \complexity O(n F) PRG expansions. n is `depth`. Each level expands every frontier node (two `expand` calls, one per share) and, while the secret path is live, one `prepare_level`. F is at most the number of distinct query leaves; the function rejects more than 2^20. Setup sorts the q query ids.
|
||||
/// \rounds No sockets. While the path is live, each level calls `prepare_level`, which samples one `ds_cw_pads` and two `ds_and_pads` and then `open_cw`.
|
||||
/// \communication none in this function.
|
||||
/// \preprocessing The `ds_randomness` tape: per live level, `ds_sample_cw` draws two 128-bit blocks and two bits, and each of the two AND pads is one `ds_sample_and`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
|
|
@ -685,6 +765,10 @@ auto geneval_sequence(InputT x0, InputT x1, ForwardIterator begin,
|
|||
std::move(qs), rng.root, rng.pad, y);
|
||||
}
|
||||
|
||||
/// \complexity O(n F) PRG expansions. n is `depth`. Each level expands every frontier node (two `expand` calls, one per share) and, while the secret path is live, one `prepare_level`. F is at most the number of distinct query leaves; the function rejects more than 2^20. Setup sorts the q query ids.
|
||||
/// \rounds No sockets. While the path is live, each level calls `prepare_level`, which samples one `ds_cw_pads` and two `ds_and_pads` and then `open_cw`.
|
||||
/// \communication none in this function.
|
||||
/// \preprocessing The `ds_randomness` tape: per live level, `ds_sample_cw` draws two 128-bit blocks and two bits, and each of the two AND pads is one `ds_sample_and`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
|
|
@ -721,6 +805,8 @@ struct geneval_cmp_result
|
|||
uint64_t addend1 = 0;
|
||||
uint64_t mask = 0;
|
||||
std::size_t live_levels = 0;
|
||||
proof_token proof0{};
|
||||
proof_token proof1{};
|
||||
};
|
||||
|
||||
/// @name Comparison geneval
|
||||
|
|
@ -762,7 +848,7 @@ geneval_cmp_result geneval_cmp(InputT x0, InputT x1,
|
|||
return out;
|
||||
|
||||
auto keys = make_dpf_doerner_shelat(std::move(x0), std::move(x1),
|
||||
std::move(rng), std::move(spec));
|
||||
std::move(rng), std::move(spec), dpf::verifiable{});
|
||||
const auto & k0 = keys.first;
|
||||
const auto & k1 = keys.second;
|
||||
using key_type = std::decay_t<decltype(k0)>;
|
||||
|
|
@ -792,11 +878,19 @@ geneval_cmp_result geneval_cmp(InputT x0, InputT x1,
|
|||
if constexpr (key_type::cmp_block == 0)
|
||||
out.value_cw[level] = k0.value_cw(level);
|
||||
}
|
||||
detail::vdpf::init_proof(out.proof0, k0);
|
||||
detail::vdpf::init_proof(out.proof1, k1);
|
||||
auto path0 = make_basic_path_memoizer(k0);
|
||||
auto path1 = make_basic_path_memoizer(k1);
|
||||
for (auto it = begin; it != end; ++it)
|
||||
{
|
||||
out.party0.push_back(eval_point(dpf::cmp, k0, *it).raw());
|
||||
out.party1.push_back(eval_point(dpf::cmp, k1, *it).raw());
|
||||
out.party0.push_back(
|
||||
detail::incr::eval_cmp_point_impl(k0, *it, path0, &out.proof0).raw());
|
||||
out.party1.push_back(
|
||||
detail::incr::eval_cmp_point_impl(k1, *it, path1, &out.proof1).raw());
|
||||
}
|
||||
detail::vdpf::fold_output_binding(out.proof0, k0);
|
||||
detail::vdpf::fold_output_binding(out.proof1, k1);
|
||||
return out;
|
||||
}
|
||||
|
||||
|
|
@ -823,7 +917,7 @@ geneval_cmp_result geneval_cmp(arith_input_t, InputT x0, InputT x1,
|
|||
return out;
|
||||
|
||||
auto keys = make_dpf_doerner_shelat(arith_input, std::move(x0), std::move(x1),
|
||||
std::move(rng), std::move(spec));
|
||||
std::move(rng), std::move(spec), dpf::verifiable{});
|
||||
const auto & k0 = keys.first;
|
||||
const auto & k1 = keys.second;
|
||||
using key_type = std::decay_t<decltype(k0)>;
|
||||
|
|
@ -853,11 +947,19 @@ geneval_cmp_result geneval_cmp(arith_input_t, InputT x0, InputT x1,
|
|||
if constexpr (key_type::cmp_block == 0)
|
||||
out.value_cw[level] = k0.value_cw(level);
|
||||
}
|
||||
detail::vdpf::init_proof(out.proof0, k0);
|
||||
detail::vdpf::init_proof(out.proof1, k1);
|
||||
auto path0 = make_basic_path_memoizer(k0);
|
||||
auto path1 = make_basic_path_memoizer(k1);
|
||||
for (auto it = begin; it != end; ++it)
|
||||
{
|
||||
out.party0.push_back(eval_point(dpf::cmp, k0, *it).raw());
|
||||
out.party1.push_back(eval_point(dpf::cmp, k1, *it).raw());
|
||||
out.party0.push_back(
|
||||
detail::incr::eval_cmp_point_impl(k0, *it, path0, &out.proof0).raw());
|
||||
out.party1.push_back(
|
||||
detail::incr::eval_cmp_point_impl(k1, *it, path1, &out.proof1).raw());
|
||||
}
|
||||
detail::vdpf::fold_output_binding(out.proof0, k0);
|
||||
detail::vdpf::fold_output_binding(out.proof1, k1);
|
||||
return out;
|
||||
}
|
||||
|
||||
|
|
@ -903,6 +1005,124 @@ geneval_cmp_result geneval_cmp(arith_input_t, InputT x0, InputT x1,
|
|||
|
||||
/// @}
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
/// @brief Copy party shares from a geneval result into caller buffers.
|
||||
template <typename Result, typename Buf0, typename Buf1>
|
||||
void geneval_fill_buffers(const Result & r, Buf0 & buf0, Buf1 & buf1)
|
||||
{
|
||||
const std::size_t n = r.party0.size();
|
||||
if (utils::size(buf0) < n || utils::size(buf1) < n)
|
||||
throw std::length_error("geneval: output buffer is too small");
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
buf0[i] = r.party0[i];
|
||||
buf1[i] = r.party1[i];
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/// @name Geneval into caller buffers
|
||||
/// @details Thin overloads that run the same trie walk, then copy party shares
|
||||
/// into `buf0` / `buf1` (same layout as `eval_interval` / `eval_sequence`
|
||||
/// output buffers). Memoizer arguments for the fused trie are internal;
|
||||
/// path memoizers live on `geneval_cmp` / `geneval_ic`.
|
||||
/// @{
|
||||
|
||||
/// \complexity O(n F) PRG expansions. n is `depth`. Each level expands every frontier node (two `expand` calls, one per share) and, while the secret path is live, one `prepare_level`. F is at most the number of distinct query leaves; the function rejects more than 2^20. Setup sorts the q query ids.
|
||||
/// \rounds No sockets. While the path is live, each level calls `prepare_level`, which samples one `ds_cw_pads` and two `ds_and_pads` and then `open_cw`.
|
||||
/// \communication none in this function.
|
||||
/// \preprocessing The `ds_randomness` tape: per live level, `ds_sample_cw` draws two 128-bit blocks and two bits, and each of the two AND pads is one `ds_sample_and`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
typename OutputT,
|
||||
typename RootSampler,
|
||||
typename PadRng,
|
||||
typename Buf0,
|
||||
typename Buf1>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto geneval_point(InputT x0, InputT x1, InputT query,
|
||||
ds_randomness<RootSampler, PadRng> rng, OutputT y, Buf0 & buf0, Buf1 & buf1)
|
||||
{
|
||||
auto r = geneval_point<InteriorPRG, ExteriorPRG>(std::move(x0),
|
||||
std::move(x1), query, std::move(rng), std::move(y));
|
||||
detail::geneval_fill_buffers(r, buf0, buf1);
|
||||
return r;
|
||||
}
|
||||
|
||||
/// \complexity O(n F) PRG expansions. n is `depth`. Each level expands every frontier node (two `expand` calls, one per share) and, while the secret path is live, one `prepare_level`. F is at most the number of distinct query leaves; the function rejects more than 2^20. Setup sorts the q query ids.
|
||||
/// \rounds No sockets. While the path is live, each level calls `prepare_level`, which samples one `ds_cw_pads` and two `ds_and_pads` and then `open_cw`.
|
||||
/// \communication none in this function.
|
||||
/// \preprocessing The `ds_randomness` tape: per live level, `ds_sample_cw` draws two 128-bit blocks and two bits, and each of the two AND pads is one `ds_sample_and`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
typename OutputT,
|
||||
typename RootSampler,
|
||||
typename PadRng,
|
||||
typename Buf0,
|
||||
typename Buf1>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto geneval_interval(InputT x0, InputT x1, InputT from, InputT to,
|
||||
ds_randomness<RootSampler, PadRng> rng, OutputT y, Buf0 & buf0, Buf1 & buf1)
|
||||
{
|
||||
auto r = geneval_interval<InteriorPRG, ExteriorPRG>(std::move(x0),
|
||||
std::move(x1), from, to, std::move(rng), std::move(y));
|
||||
detail::geneval_fill_buffers(r, buf0, buf1);
|
||||
return r;
|
||||
}
|
||||
|
||||
/// \complexity O(n F) PRG expansions. n is `depth`. Each level expands every frontier node (two `expand` calls, one per share) and, while the secret path is live, one `prepare_level`. F is at most the number of distinct query leaves; the function rejects more than 2^20. Setup sorts the q query ids.
|
||||
/// \rounds No sockets. While the path is live, each level calls `prepare_level`, which samples one `ds_cw_pads` and two `ds_and_pads` and then `open_cw`.
|
||||
/// \communication none in this function.
|
||||
/// \preprocessing The `ds_randomness` tape: per live level, `ds_sample_cw` draws two 128-bit blocks and two bits, and each of the two AND pads is one `ds_sample_and`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
typename OutputT,
|
||||
typename RootSampler,
|
||||
typename PadRng,
|
||||
typename Buf0,
|
||||
typename Buf1>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto geneval_full(InputT x0, InputT x1,
|
||||
ds_randomness<RootSampler, PadRng> rng, OutputT y, Buf0 & buf0, Buf1 & buf1)
|
||||
{
|
||||
auto r = geneval_full<InteriorPRG, ExteriorPRG>(std::move(x0),
|
||||
std::move(x1), std::move(rng), std::move(y));
|
||||
detail::geneval_fill_buffers(r, buf0, buf1);
|
||||
return r;
|
||||
}
|
||||
|
||||
/// \complexity O(n F) PRG expansions. n is `depth`. Each level expands every frontier node (two `expand` calls, one per share) and, while the secret path is live, one `prepare_level`. F is at most the number of distinct query leaves; the function rejects more than 2^20. Setup sorts the q query ids.
|
||||
/// \rounds No sockets. While the path is live, each level calls `prepare_level`, which samples one `ds_cw_pads` and two `ds_and_pads` and then `open_cw`.
|
||||
/// \communication none in this function.
|
||||
/// \preprocessing The `ds_randomness` tape: per live level, `ds_sample_cw` draws two 128-bit blocks and two bits, and each of the two AND pads is one `ds_sample_and`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
typename OutputT,
|
||||
typename RootSampler,
|
||||
typename PadRng,
|
||||
typename ForwardIterator,
|
||||
typename Buf0,
|
||||
typename Buf1>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto geneval_sequence(InputT x0, InputT x1, ForwardIterator begin,
|
||||
ForwardIterator end, ds_randomness<RootSampler, PadRng> rng, OutputT y,
|
||||
Buf0 & buf0, Buf1 & buf1)
|
||||
{
|
||||
auto r = geneval_sequence<InteriorPRG, ExteriorPRG>(std::move(x0),
|
||||
std::move(x1), begin, end, std::move(rng), std::move(y));
|
||||
detail::geneval_fill_buffers(r, buf0, buf1);
|
||||
return r;
|
||||
}
|
||||
|
||||
/// @}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_GENEVAL_HPP__
|
||||
|
|
|
|||
924
include/dpf/gf2.hpp
Normal file
924
include/dpf/gf2.hpp
Normal file
|
|
@ -0,0 +1,924 @@
|
|||
/// @file dpf/gf2.hpp
|
||||
/// @brief Characteristic-2 field elements as DPF output types.
|
||||
/// @details `dpf::gf2`, `gf22`, `gf24`, `gf28`, `gf216`, `gf232`, and `gf264`
|
||||
/// are GF(2^k) for k = 1, 2, 4, 8, 16, 32, 64. Addition and
|
||||
/// subtraction are XOR. Multiplication is the polynomial product
|
||||
/// in the standard basis (low bit is the coefficient of 1).
|
||||
///
|
||||
/// The moduli match laneint on mocha2. Widths 1, 2, 4, and 16 are
|
||||
/// the sparse irreducibles in `laneint/gf2_poly_basis.hpp` and the
|
||||
/// epi multipliers. GF(2^8) is the AES polynomial that
|
||||
/// `gf256_mul_u8` uses (`0x11B`), not the sparse `0x11D` alternative.
|
||||
/// GF(2^32) and GF(2^64) use the irreducible basis polynomials.
|
||||
/// The epi “default” reducers `x^32+x^7+1` and `x^64+x^4+1` factor,
|
||||
/// so they are not the field. Leaf scaling of `gf28` and `gf216`
|
||||
/// is the FAST'13 nibble `pshufb` (Plank, Greenan, Miller): one
|
||||
/// constant builds 16-byte tables, and every lane is a shuffle.
|
||||
/// `detail::shamir_field` inverts a
|
||||
/// nonzero element by Fermat, `a^{2^k - 2}`. Shamir points are the
|
||||
/// integers `1 .. N` as bit patterns, so `N` must be less than `2^k`
|
||||
/// or two parties land on the same element.
|
||||
/// @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_GF2_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_GF2_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <ostream>
|
||||
#include <stdexcept>
|
||||
#include <type_traits>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
#include "simde/simde/x86/avx2.h"
|
||||
|
||||
#include "dpf/leaf_arithmetic.hpp"
|
||||
#include "dpf/utils.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace gf2_detail
|
||||
{
|
||||
|
||||
template <unsigned Bits>
|
||||
struct field;
|
||||
|
||||
template <>
|
||||
struct field<1>
|
||||
{
|
||||
using word = std::uint8_t;
|
||||
using wide = std::uint16_t;
|
||||
static constexpr unsigned bits = 1;
|
||||
/// `x + 1`
|
||||
static constexpr wide modulus = 0x3u;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct field<2>
|
||||
{
|
||||
using word = std::uint8_t;
|
||||
using wide = std::uint16_t;
|
||||
static constexpr unsigned bits = 2;
|
||||
/// `x^2 + x + 1`
|
||||
static constexpr wide modulus = 0x7u;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct field<4>
|
||||
{
|
||||
using word = std::uint8_t;
|
||||
using wide = std::uint16_t;
|
||||
static constexpr unsigned bits = 4;
|
||||
/// `x^4 + x + 1`
|
||||
static constexpr wide modulus = 0x13u;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct field<8>
|
||||
{
|
||||
using word = std::uint8_t;
|
||||
using wide = std::uint16_t;
|
||||
static constexpr unsigned bits = 8;
|
||||
/// AES: `x^8 + x^4 + x^3 + x + 1`
|
||||
static constexpr wide modulus = 0x11bu;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct field<16>
|
||||
{
|
||||
using word = std::uint16_t;
|
||||
using wide = std::uint32_t;
|
||||
static constexpr unsigned bits = 16;
|
||||
/// `x^16 + x^5 + x^3 + x^2 + 1`
|
||||
static constexpr wide modulus = 0x1002du;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct field<32>
|
||||
{
|
||||
using word = std::uint32_t;
|
||||
using wide = std::uint64_t;
|
||||
static constexpr unsigned bits = 32;
|
||||
/// `x^32 + x^31 + x^28 + x^21 + 1`
|
||||
static constexpr wide modulus = 0x190200001ull;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct field<64>
|
||||
{
|
||||
using word = std::uint64_t;
|
||||
using wide = unsigned __int128;
|
||||
static constexpr unsigned bits = 64;
|
||||
/// `x^64 + x^63 + x^62 + x^53 + 1`
|
||||
static constexpr wide modulus =
|
||||
(wide{1} << 64) | (wide{1} << 63) | (wide{1} << 62) | (wide{1} << 53) | wide{1};
|
||||
};
|
||||
|
||||
template <unsigned Bits>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_CONST
|
||||
constexpr typename field<Bits>::word mask() noexcept
|
||||
{
|
||||
using word = typename field<Bits>::word;
|
||||
if constexpr (Bits >= sizeof(word) * 8u)
|
||||
return static_cast<word>(~word{0});
|
||||
else
|
||||
return static_cast<word>((word{1} << Bits) - word{1});
|
||||
}
|
||||
|
||||
template <unsigned Bits>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_CONST
|
||||
constexpr typename field<Bits>::word reduce(typename field<Bits>::wide p) noexcept
|
||||
{
|
||||
using wide = typename field<Bits>::wide;
|
||||
constexpr wide mod = field<Bits>::modulus;
|
||||
for (int bit = static_cast<int>(2u * Bits) - 2; bit >= static_cast<int>(Bits); --bit)
|
||||
{
|
||||
if (((p >> bit) & wide{1}) != 0)
|
||||
p ^= mod << (bit - static_cast<int>(Bits));
|
||||
}
|
||||
return static_cast<typename field<Bits>::word>(p);
|
||||
}
|
||||
|
||||
template <unsigned Bits>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_CONST
|
||||
constexpr typename field<Bits>::word mul(typename field<Bits>::word a,
|
||||
typename field<Bits>::word b) noexcept
|
||||
{
|
||||
using wide = typename field<Bits>::wide;
|
||||
a = static_cast<typename field<Bits>::word>(a & mask<Bits>());
|
||||
b = static_cast<typename field<Bits>::word>(b & mask<Bits>());
|
||||
wide p = 0;
|
||||
for (unsigned i = 0; i < Bits; ++i)
|
||||
{
|
||||
if (((b >> i) & 1u) != 0u)
|
||||
p ^= static_cast<wide>(a) << i;
|
||||
}
|
||||
return reduce<Bits>(p);
|
||||
}
|
||||
|
||||
/// @brief Multiplicative inverse, `a^{2^Bits - 2}`.
|
||||
/// @param a a nonzero field element
|
||||
/// @return `a^{-1}`
|
||||
/// @throws std::invalid_argument if `a` is zero
|
||||
template <unsigned Bits>
|
||||
typename field<Bits>::word inv(typename field<Bits>::word a)
|
||||
{
|
||||
using word = typename field<Bits>::word;
|
||||
a = static_cast<word>(a & mask<Bits>());
|
||||
if (a == 0)
|
||||
throw std::invalid_argument("shamir: inverse of zero");
|
||||
// Exponent 2^Bits - 2 is Bits-1 ones followed by a zero.
|
||||
word result{1};
|
||||
for (int bit = static_cast<int>(Bits) - 1; bit >= 0; --bit)
|
||||
{
|
||||
result = mul<Bits>(result, result);
|
||||
if (bit != 0)
|
||||
result = mul<Bits>(result, a);
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
/// @brief Multiply by `x` in AES GF(2^8), modulus `0x11B`.
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr std::uint8_t gf28_mul_x(std::uint8_t a) noexcept
|
||||
{
|
||||
const std::uint8_t hi = static_cast<std::uint8_t>(a & 0x80u);
|
||||
a = static_cast<std::uint8_t>(static_cast<std::uint8_t>(a << 1) ^ (hi != 0 ? 0x1bu : 0));
|
||||
return a;
|
||||
}
|
||||
|
||||
/// @brief Multiply by `x` in GF(2^16), modulus `x^16+x^5+x^3+x^2+1`.
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr std::uint16_t gf216_mul_x(std::uint16_t a) noexcept
|
||||
{
|
||||
const std::uint16_t hi = static_cast<std::uint16_t>(a & 0x8000u);
|
||||
a = static_cast<std::uint16_t>(static_cast<std::uint16_t>(a << 1) ^ (hi != 0 ? 0x2du : 0));
|
||||
return a;
|
||||
}
|
||||
|
||||
/// @brief 16 products of `base` with a nibble, built from four multiplies by `x`.
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void nibble_products_u8(std::uint8_t base, std::uint8_t out[16]) noexcept
|
||||
{
|
||||
const std::uint8_t b0 = base;
|
||||
const std::uint8_t b1 = gf28_mul_x(b0);
|
||||
const std::uint8_t b2 = gf28_mul_x(b1);
|
||||
const std::uint8_t b3 = gf28_mul_x(b2);
|
||||
out[0] = 0;
|
||||
for (unsigned n = 1; n < 16u; ++n)
|
||||
{
|
||||
std::uint8_t r = 0;
|
||||
if ((n & 1u) != 0) r = static_cast<std::uint8_t>(r ^ b0);
|
||||
if ((n & 2u) != 0) r = static_cast<std::uint8_t>(r ^ b1);
|
||||
if ((n & 4u) != 0) r = static_cast<std::uint8_t>(r ^ b2);
|
||||
if ((n & 8u) != 0) r = static_cast<std::uint8_t>(r ^ b3);
|
||||
out[n] = r;
|
||||
}
|
||||
}
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void nibble_products_u16(std::uint16_t base, std::uint8_t lo[16], std::uint8_t hi[16]) noexcept
|
||||
{
|
||||
const std::uint16_t b0 = base;
|
||||
const std::uint16_t b1 = gf216_mul_x(b0);
|
||||
const std::uint16_t b2 = gf216_mul_x(b1);
|
||||
const std::uint16_t b3 = gf216_mul_x(b2);
|
||||
lo[0] = 0;
|
||||
hi[0] = 0;
|
||||
for (unsigned n = 1; n < 16u; ++n)
|
||||
{
|
||||
std::uint16_t r = 0;
|
||||
if ((n & 1u) != 0) r = static_cast<std::uint16_t>(r ^ b0);
|
||||
if ((n & 2u) != 0) r = static_cast<std::uint16_t>(r ^ b1);
|
||||
if ((n & 4u) != 0) r = static_cast<std::uint16_t>(r ^ b2);
|
||||
if ((n & 8u) != 0) r = static_cast<std::uint16_t>(r ^ b3);
|
||||
lo[n] = static_cast<std::uint8_t>(r);
|
||||
hi[n] = static_cast<std::uint8_t>(r >> 8);
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief FAST'13 nibble tables for one GF(2^8) constant.
|
||||
/// @details Plank, Greenan, Miller, "Screaming Fast Galois Field Arithmetic
|
||||
/// Using Intel SIMD Instructions". `y * a = (y * a_lo) XOR (y * (a_hi << 4))`,
|
||||
/// each half a 16-entry `pshufb`.
|
||||
struct gf28_shufb
|
||||
{
|
||||
simde__m128i lo;
|
||||
simde__m128i hi;
|
||||
};
|
||||
|
||||
inline const gf28_shufb & gf28_shufb_for(std::uint8_t y) noexcept
|
||||
{
|
||||
struct tables
|
||||
{
|
||||
gf28_shufb t[256];
|
||||
tables() noexcept
|
||||
{
|
||||
for (unsigned yy = 0; yy < 256u; ++yy)
|
||||
{
|
||||
alignas(16) std::uint8_t lo[16];
|
||||
alignas(16) std::uint8_t hi[16];
|
||||
nibble_products_u8(static_cast<std::uint8_t>(yy), lo);
|
||||
const std::uint8_t x4 = gf28_mul_x(gf28_mul_x(gf28_mul_x(gf28_mul_x(
|
||||
static_cast<std::uint8_t>(yy)))));
|
||||
nibble_products_u8(x4, hi);
|
||||
t[yy].lo = simde_mm_load_si128(reinterpret_cast<const simde__m128i *>(lo));
|
||||
t[yy].hi = simde_mm_load_si128(reinterpret_cast<const simde__m128i *>(hi));
|
||||
}
|
||||
}
|
||||
};
|
||||
static const tables all;
|
||||
return all.t[y];
|
||||
}
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
simde__m128i gf28_scale_m128(simde__m128i lut_lo, simde__m128i lut_hi, simde__m128i a) noexcept
|
||||
{
|
||||
const simde__m128i m0f = simde_mm_set1_epi8(0x0f);
|
||||
const simde__m128i lo = simde_mm_and_si128(a, m0f);
|
||||
const simde__m128i hi = simde_mm_and_si128(simde_mm_srli_epi16(a, 4), m0f);
|
||||
return simde_mm_xor_si128(simde_mm_shuffle_epi8(lut_lo, lo),
|
||||
simde_mm_shuffle_epi8(lut_hi, hi));
|
||||
}
|
||||
|
||||
/// @brief Four nibble positions of one GF(2^16) constant. Each position is a
|
||||
/// low-byte table and a high-byte table.
|
||||
struct gf216_shufb
|
||||
{
|
||||
simde__m128i lo[4];
|
||||
simde__m128i hi[4];
|
||||
};
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
gf216_shufb gf216_shufb_prepare(std::uint16_t k) noexcept
|
||||
{
|
||||
gf216_shufb L{};
|
||||
std::uint16_t base = k;
|
||||
for (unsigned nib = 0; nib < 4u; ++nib)
|
||||
{
|
||||
alignas(16) std::uint8_t lo[16];
|
||||
alignas(16) std::uint8_t hi[16];
|
||||
nibble_products_u16(base, lo, hi);
|
||||
L.lo[nib] = simde_mm_load_si128(reinterpret_cast<const simde__m128i *>(lo));
|
||||
L.hi[nib] = simde_mm_load_si128(reinterpret_cast<const simde__m128i *>(hi));
|
||||
base = gf216_mul_x(gf216_mul_x(gf216_mul_x(gf216_mul_x(base))));
|
||||
}
|
||||
return L;
|
||||
}
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
simde__m128i gf216_scale_m128(const gf216_shufb & L, simde__m128i x) noexcept
|
||||
{
|
||||
const simde__m128i m4 = simde_mm_set1_epi16(0x000f);
|
||||
const simde__m128i pick_ev = simde_mm_setr_epi8(
|
||||
0, static_cast<int8_t>(0x80), 2, static_cast<int8_t>(0x80),
|
||||
4, static_cast<int8_t>(0x80), 6, static_cast<int8_t>(0x80),
|
||||
8, static_cast<int8_t>(0x80), 10, static_cast<int8_t>(0x80),
|
||||
12, static_cast<int8_t>(0x80), 14, static_cast<int8_t>(0x80));
|
||||
const simde__m128i ix0 = simde_mm_shuffle_epi8(simde_mm_and_si128(x, m4), pick_ev);
|
||||
const simde__m128i ix1 = simde_mm_shuffle_epi8(
|
||||
simde_mm_and_si128(simde_mm_srli_epi16(x, 4), m4), pick_ev);
|
||||
const simde__m128i ix2 = simde_mm_shuffle_epi8(
|
||||
simde_mm_and_si128(simde_mm_srli_epi16(x, 8), m4), pick_ev);
|
||||
const simde__m128i ix3 = simde_mm_shuffle_epi8(
|
||||
simde_mm_and_si128(simde_mm_srli_epi16(x, 12), m4), pick_ev);
|
||||
auto part = [](simde__m128i tl, simde__m128i th, simde__m128i ix) noexcept {
|
||||
const simde__m128i pl = simde_mm_shuffle_epi8(tl, ix);
|
||||
const simde__m128i ph = simde_mm_shuffle_epi8(th, ix);
|
||||
return simde_mm_or_si128(pl, simde_mm_slli_epi16(ph, 8));
|
||||
};
|
||||
return simde_mm_xor_si128(
|
||||
simde_mm_xor_si128(part(L.lo[0], L.hi[0], ix0), part(L.lo[1], L.hi[1], ix1)),
|
||||
simde_mm_xor_si128(part(L.lo[2], L.hi[2], ix2), part(L.lo[3], L.hi[3], ix3)));
|
||||
}
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void scale_gf28(unsigned char * dst, const unsigned char * src, std::size_t n,
|
||||
std::uint8_t k) noexcept
|
||||
{
|
||||
if (k == 0)
|
||||
{
|
||||
std::memset(dst, 0, n);
|
||||
return;
|
||||
}
|
||||
if (k == 1)
|
||||
{
|
||||
if (dst != src)
|
||||
std::memcpy(dst, src, n);
|
||||
return;
|
||||
}
|
||||
const gf28_shufb & lut = gf28_shufb_for(k);
|
||||
const simde__m256i lo256 = simde_mm256_broadcastsi128_si256(lut.lo);
|
||||
const simde__m256i hi256 = simde_mm256_broadcastsi128_si256(lut.hi);
|
||||
const simde__m256i m0f = simde_mm256_set1_epi8(0x0f);
|
||||
std::size_t i = 0;
|
||||
for (; i + 32 <= n; i += 32)
|
||||
{
|
||||
const simde__m256i a = simde_mm256_loadu_si256(src + i);
|
||||
const simde__m256i lo = simde_mm256_and_si256(a, m0f);
|
||||
const simde__m256i hi = simde_mm256_and_si256(simde_mm256_srli_epi16(a, 4), m0f);
|
||||
const simde__m256i r = simde_mm256_xor_si256(
|
||||
simde_mm256_shuffle_epi8(lo256, lo),
|
||||
simde_mm256_shuffle_epi8(hi256, hi));
|
||||
simde_mm256_storeu_si256(dst + i, r);
|
||||
}
|
||||
for (; i + 16 <= n; i += 16)
|
||||
{
|
||||
const simde__m128i a = simde_mm_loadu_si128(src + i);
|
||||
simde_mm_storeu_si128(dst + i, gf28_scale_m128(lut.lo, lut.hi, a));
|
||||
}
|
||||
for (; i < n; ++i)
|
||||
dst[i] = mul<8>(src[i], k);
|
||||
}
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void scale_gf216(unsigned char * dst, const unsigned char * src, std::size_t n,
|
||||
std::uint16_t k) noexcept
|
||||
{
|
||||
if (k == 0)
|
||||
{
|
||||
std::memset(dst, 0, n);
|
||||
return;
|
||||
}
|
||||
if (k == 1)
|
||||
{
|
||||
if (dst != src)
|
||||
std::memcpy(dst, src, n);
|
||||
return;
|
||||
}
|
||||
const gf216_shufb L = gf216_shufb_prepare(k);
|
||||
std::size_t i = 0;
|
||||
for (; i + 16 <= n; i += 16)
|
||||
{
|
||||
const simde__m128i a = simde_mm_loadu_si128(src + i);
|
||||
simde_mm_storeu_si128(dst + i, gf216_scale_m128(L, a));
|
||||
}
|
||||
for (; i + sizeof(std::uint16_t) <= n; i += sizeof(std::uint16_t))
|
||||
{
|
||||
std::uint16_t a{};
|
||||
std::memcpy(&a, src + i, sizeof(a));
|
||||
const std::uint16_t c = mul<16>(a, k);
|
||||
std::memcpy(dst + i, &c, sizeof(c));
|
||||
}
|
||||
for (; i < n; ++i)
|
||||
dst[i] = 0;
|
||||
}
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void xor_bytes(unsigned char * dst, const unsigned char * a,
|
||||
const unsigned char * b, std::size_t n) noexcept
|
||||
{
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
dst[i] = static_cast<unsigned char>(a[i] ^ b[i]);
|
||||
}
|
||||
|
||||
template <unsigned Bits>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void scale_bytes(unsigned char * dst, const unsigned char * src, std::size_t n,
|
||||
typename field<Bits>::word k) noexcept
|
||||
{
|
||||
using word = typename field<Bits>::word;
|
||||
k = static_cast<word>(k & mask<Bits>());
|
||||
if constexpr (Bits == 1)
|
||||
{
|
||||
const unsigned char fill = (k & 1u) ? static_cast<unsigned char>(0xffu) : 0;
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
dst[i] = static_cast<unsigned char>(src[i] & fill);
|
||||
}
|
||||
else if constexpr (Bits < 8)
|
||||
{
|
||||
constexpr unsigned per_byte = 8u / Bits;
|
||||
constexpr unsigned lane_mask = (1u << Bits) - 1u;
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
unsigned out = 0;
|
||||
const unsigned in = src[i];
|
||||
for (unsigned lane = 0; lane < per_byte; ++lane)
|
||||
{
|
||||
const auto a = static_cast<word>((in >> (lane * Bits)) & lane_mask);
|
||||
out |= static_cast<unsigned>(mul<Bits>(a, k)) << (lane * Bits);
|
||||
}
|
||||
dst[i] = static_cast<unsigned char>(out);
|
||||
}
|
||||
}
|
||||
else if constexpr (Bits == 8)
|
||||
{
|
||||
scale_gf28(dst, src, n, static_cast<std::uint8_t>(k));
|
||||
}
|
||||
else if constexpr (Bits == 16)
|
||||
{
|
||||
scale_gf216(dst, src, n, static_cast<std::uint16_t>(k));
|
||||
}
|
||||
else
|
||||
{
|
||||
std::size_t i = 0;
|
||||
for (; i + sizeof(word) <= n; i += sizeof(word))
|
||||
{
|
||||
word a{};
|
||||
std::memcpy(&a, src + i, sizeof(word));
|
||||
const word c = mul<Bits>(a, k);
|
||||
std::memcpy(dst + i, &c, sizeof(word));
|
||||
}
|
||||
for (; i < n; ++i)
|
||||
dst[i] = 0;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename NodeT>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
NodeT xor_node(const NodeT & a, const NodeT & b) noexcept
|
||||
{
|
||||
unsigned char aa[sizeof(NodeT)];
|
||||
unsigned char bb[sizeof(NodeT)];
|
||||
unsigned char cc[sizeof(NodeT)];
|
||||
std::memcpy(aa, std::addressof(a), sizeof(NodeT));
|
||||
std::memcpy(bb, std::addressof(b), sizeof(NodeT));
|
||||
xor_bytes(cc, aa, bb, sizeof(NodeT));
|
||||
NodeT out;
|
||||
std::memcpy(std::addressof(out), cc, sizeof(NodeT));
|
||||
return out;
|
||||
}
|
||||
|
||||
template <unsigned Bits, typename NodeT>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
NodeT scale_node(const NodeT & a, typename field<Bits>::word k) noexcept
|
||||
{
|
||||
unsigned char src[sizeof(NodeT)];
|
||||
unsigned char dst[sizeof(NodeT)];
|
||||
std::memcpy(src, std::addressof(a), sizeof(NodeT));
|
||||
scale_bytes<Bits>(dst, src, sizeof(NodeT), k);
|
||||
NodeT out;
|
||||
std::memcpy(std::addressof(out), dst, sizeof(NodeT));
|
||||
return out;
|
||||
}
|
||||
|
||||
} // namespace gf2_detail
|
||||
|
||||
/// @brief Element of GF(2^Bits) in the standard polynomial basis.
|
||||
/// @tparam Bits field degree. One of 1, 2, 4, 8, 16, 32, 64.
|
||||
template <unsigned Bits>
|
||||
class gf2n
|
||||
{
|
||||
public:
|
||||
using integral_type = typename gf2_detail::field<Bits>::word;
|
||||
static constexpr unsigned bits = Bits;
|
||||
static constexpr bool dpf_gf2n = true;
|
||||
static constexpr bool dpf_point_group = true;
|
||||
|
||||
/// @brief The zero element.
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr gf2n() noexcept = default;
|
||||
|
||||
/// @brief Bit embedding of an integer. A negative value is the field
|
||||
/// negation of its magnitude, which is the same element: every
|
||||
/// element is its own additive inverse.
|
||||
/// @tparam T integral type
|
||||
/// @param v the integer to embed
|
||||
template <typename T, typename = std::enable_if_t<std::is_integral_v<T>>>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr gf2n(T v) noexcept
|
||||
{
|
||||
if constexpr (std::is_signed_v<T>)
|
||||
{
|
||||
if (v < 0)
|
||||
{
|
||||
using unsigned_type = std::make_unsigned_t<T>;
|
||||
const auto mag = static_cast<unsigned_type>(0)
|
||||
- static_cast<unsigned_type>(v);
|
||||
val = static_cast<integral_type>(mag) & gf2_detail::mask<Bits>();
|
||||
}
|
||||
else
|
||||
{
|
||||
val = static_cast<integral_type>(v) & gf2_detail::mask<Bits>();
|
||||
}
|
||||
}
|
||||
else
|
||||
{
|
||||
val = static_cast<integral_type>(v) & gf2_detail::mask<Bits>();
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Low `Bits` of a PRG block. Every bit string is a field element.
|
||||
/// @param bytes the PRG output
|
||||
/// @param n the number of bytes available
|
||||
/// @return the field element
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
static gf2n from_seed(const void * bytes, std::size_t n) noexcept
|
||||
{
|
||||
integral_type w{};
|
||||
if (bytes != nullptr && n != 0)
|
||||
{
|
||||
const std::size_t take = n < sizeof(w) ? n : sizeof(w);
|
||||
std::memcpy(&w, bytes, take);
|
||||
}
|
||||
return gf2n{w};
|
||||
}
|
||||
|
||||
/// @brief Reduced representative in `0 .. 2^Bits - 1`.
|
||||
/// @return the stored field element
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_PURE
|
||||
constexpr integral_type raw() const noexcept
|
||||
{
|
||||
return static_cast<integral_type>(val & gf2_detail::mask<Bits>());
|
||||
}
|
||||
|
||||
/// @brief Same value as `raw()`.
|
||||
/// @return the stored field element
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_PURE
|
||||
explicit constexpr operator integral_type() const noexcept { return raw(); }
|
||||
|
||||
private:
|
||||
integral_type val{};
|
||||
};
|
||||
|
||||
/// @brief GF(2).
|
||||
using gf2 = gf2n<1>;
|
||||
/// @brief GF(4) = GF(2^2), modulus `x^2 + x + 1`.
|
||||
using gf22 = gf2n<2>;
|
||||
/// @brief GF(16) = GF(2^4), modulus `x^4 + x + 1`.
|
||||
using gf24 = gf2n<4>;
|
||||
/// @brief GF(256) = GF(2^8), AES modulus `x^8 + x^4 + x^3 + x + 1`.
|
||||
using gf28 = gf2n<8>;
|
||||
/// @brief GF(2^16), modulus `x^16 + x^5 + x^3 + x^2 + 1`.
|
||||
using gf216 = gf2n<16>;
|
||||
/// @brief GF(2^32), modulus `x^32 + x^31 + x^28 + x^21 + 1`.
|
||||
using gf232 = gf2n<32>;
|
||||
/// @brief GF(2^64), modulus `x^64 + x^63 + x^62 + x^53 + 1`.
|
||||
using gf264 = gf2n<64>;
|
||||
|
||||
/// @brief Field addition. XOR of the reduced representatives.
|
||||
template <unsigned Bits>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr gf2n<Bits> operator+(gf2n<Bits> a, gf2n<Bits> b) noexcept
|
||||
{
|
||||
using word = typename gf2n<Bits>::integral_type;
|
||||
return gf2n<Bits>{static_cast<word>(a.raw() ^ b.raw())};
|
||||
}
|
||||
|
||||
/// @brief Field subtraction. Identical to addition.
|
||||
template <unsigned Bits>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr gf2n<Bits> operator-(gf2n<Bits> a, gf2n<Bits> b) noexcept
|
||||
{
|
||||
return a + b;
|
||||
}
|
||||
|
||||
/// @brief Additive inverse. Identical to `a` in characteristic 2.
|
||||
template <unsigned Bits>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr gf2n<Bits> operator-(gf2n<Bits> a) noexcept
|
||||
{
|
||||
return a;
|
||||
}
|
||||
|
||||
/// @brief Field XOR. Identical to addition.
|
||||
template <unsigned Bits>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr gf2n<Bits> operator^(gf2n<Bits> a, gf2n<Bits> b) noexcept
|
||||
{
|
||||
return a + b;
|
||||
}
|
||||
|
||||
/// @brief Field multiplication.
|
||||
template <unsigned Bits>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr gf2n<Bits> operator*(gf2n<Bits> a, gf2n<Bits> b) noexcept
|
||||
{
|
||||
return gf2n<Bits>{gf2_detail::mul<Bits>(a.raw(), b.raw())};
|
||||
}
|
||||
|
||||
/// @brief Equality of the reduced representatives.
|
||||
template <unsigned Bits>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr bool operator==(gf2n<Bits> a, gf2n<Bits> b) noexcept
|
||||
{
|
||||
return a.raw() == b.raw();
|
||||
}
|
||||
|
||||
/// @brief Inequality of the reduced representatives.
|
||||
template <unsigned Bits>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr bool operator!=(gf2n<Bits> a, gf2n<Bits> b) noexcept
|
||||
{
|
||||
return a.raw() != b.raw();
|
||||
}
|
||||
|
||||
/// @brief Write the reduced representative in decimal.
|
||||
template <unsigned Bits>
|
||||
std::ostream & operator<<(std::ostream & os, gf2n<Bits> a)
|
||||
{
|
||||
return os << static_cast<unsigned long long>(a.raw());
|
||||
}
|
||||
|
||||
/// @brief Non-template operators so a packed-lane proxy converts to the field
|
||||
/// and finds `+` / `-` the same way `nyble` and `twobit` do.
|
||||
#define DPF_GF_LANE_OPS(Type) \
|
||||
HEDLEY_ALWAYS_INLINE constexpr Type operator+(Type a, Type b) noexcept \
|
||||
{ \
|
||||
using word = typename Type::integral_type; \
|
||||
return Type{static_cast<word>(a.raw() ^ b.raw())}; \
|
||||
} \
|
||||
HEDLEY_ALWAYS_INLINE constexpr Type operator-(Type a, Type b) noexcept \
|
||||
{ \
|
||||
return a + b; \
|
||||
} \
|
||||
HEDLEY_ALWAYS_INLINE constexpr Type operator-(Type a) noexcept { return a; } \
|
||||
HEDLEY_ALWAYS_INLINE constexpr Type operator^(Type a, Type b) noexcept \
|
||||
{ \
|
||||
return a + b; \
|
||||
} \
|
||||
HEDLEY_ALWAYS_INLINE constexpr Type operator*(Type a, Type b) noexcept \
|
||||
{ \
|
||||
return Type{gf2_detail::mul<Type::bits>(a.raw(), b.raw())}; \
|
||||
} \
|
||||
HEDLEY_ALWAYS_INLINE constexpr bool operator==(Type a, Type b) noexcept \
|
||||
{ \
|
||||
return a.raw() == b.raw(); \
|
||||
} \
|
||||
HEDLEY_ALWAYS_INLINE constexpr bool operator!=(Type a, Type b) noexcept \
|
||||
{ \
|
||||
return !(a == b); \
|
||||
} \
|
||||
inline std::ostream & operator<<(std::ostream & os, Type a) \
|
||||
{ \
|
||||
return os << static_cast<unsigned long long>(a.raw()); \
|
||||
}
|
||||
|
||||
DPF_GF_LANE_OPS(gf2)
|
||||
DPF_GF_LANE_OPS(gf22)
|
||||
DPF_GF_LANE_OPS(gf24)
|
||||
DPF_GF_LANE_OPS(gf28)
|
||||
DPF_GF_LANE_OPS(gf216)
|
||||
DPF_GF_LANE_OPS(gf232)
|
||||
DPF_GF_LANE_OPS(gf264)
|
||||
#undef DPF_GF_LANE_OPS
|
||||
|
||||
namespace utils
|
||||
{
|
||||
|
||||
template <unsigned Bits>
|
||||
struct bitlength_of<gf2n<Bits>>
|
||||
: std::integral_constant<std::size_t, Bits>
|
||||
{ };
|
||||
|
||||
template <unsigned Bits, typename NodeT>
|
||||
struct bitlength_of_output<gf2n<Bits>, NodeT>
|
||||
: std::integral_constant<std::size_t, Bits>
|
||||
{ };
|
||||
|
||||
template <unsigned Bits>
|
||||
struct is_packed_subbyte<gf2n<Bits>> : std::bool_constant<(Bits < 8)> {};
|
||||
|
||||
template <unsigned Bits>
|
||||
struct packed_lane_bits<gf2n<Bits>>
|
||||
: std::integral_constant<std::size_t, (Bits < 8 ? Bits : 0)> {};
|
||||
|
||||
template <unsigned Bits>
|
||||
struct has_characteristic_two<gf2n<Bits>> : std::true_type {};
|
||||
|
||||
} // namespace utils
|
||||
|
||||
namespace leaf_arithmetic
|
||||
{
|
||||
|
||||
template <unsigned Bits, typename NodeT>
|
||||
struct add_t<gf2n<Bits>, NodeT, std::enable_if_t<!std::is_void_v<NodeT>>>
|
||||
{
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto operator()(const NodeT & a, const NodeT & b) const noexcept
|
||||
{
|
||||
return gf2_detail::xor_node(a, b);
|
||||
}
|
||||
};
|
||||
|
||||
template <unsigned Bits, typename NodeT>
|
||||
struct subtract_t<gf2n<Bits>, NodeT, std::enable_if_t<!std::is_void_v<NodeT>>>
|
||||
{
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto operator()(const NodeT & a, const NodeT & b) const noexcept
|
||||
{
|
||||
return gf2_detail::xor_node(a, b);
|
||||
}
|
||||
};
|
||||
|
||||
template <unsigned Bits, typename NodeT>
|
||||
struct multiply_t<gf2n<Bits>, NodeT, std::enable_if_t<!std::is_void_v<NodeT>>>
|
||||
{
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto operator()(const NodeT & a, gf2n<Bits> b) const noexcept
|
||||
{
|
||||
return gf2_detail::scale_node<Bits>(a, b.raw());
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace leaf_arithmetic
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#include "dpf/output_buffer.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
template <>
|
||||
class output_buffer<gf2> : public dynamic_packed_array<gf2>
|
||||
{
|
||||
using base = dynamic_packed_array<gf2>;
|
||||
public:
|
||||
using size_type = typename base::size_type;
|
||||
explicit output_buffer(size_type size) : base(size) { }
|
||||
output_buffer(output_buffer &&) noexcept = default;
|
||||
output_buffer(const output_buffer &) = delete;
|
||||
output_buffer & operator=(output_buffer &&) noexcept = default;
|
||||
output_buffer & operator=(const output_buffer &) = delete;
|
||||
~output_buffer() noexcept = default;
|
||||
};
|
||||
|
||||
template <>
|
||||
class output_buffer<gf22> : public dynamic_packed_array<gf22>
|
||||
{
|
||||
using base = dynamic_packed_array<gf22>;
|
||||
public:
|
||||
using size_type = typename base::size_type;
|
||||
explicit output_buffer(size_type size) : base(size) { }
|
||||
output_buffer(output_buffer &&) noexcept = default;
|
||||
output_buffer(const output_buffer &) = delete;
|
||||
output_buffer & operator=(output_buffer &&) noexcept = default;
|
||||
output_buffer & operator=(const output_buffer &) = delete;
|
||||
~output_buffer() noexcept = default;
|
||||
};
|
||||
|
||||
template <>
|
||||
class output_buffer<gf24> : public dynamic_packed_array<gf24>
|
||||
{
|
||||
using base = dynamic_packed_array<gf24>;
|
||||
public:
|
||||
using size_type = typename base::size_type;
|
||||
explicit output_buffer(size_type size) : base(size) { }
|
||||
output_buffer(output_buffer &&) noexcept = default;
|
||||
output_buffer(const output_buffer &) = delete;
|
||||
output_buffer & operator=(output_buffer &&) noexcept = default;
|
||||
output_buffer & operator=(const output_buffer &) = delete;
|
||||
~output_buffer() noexcept = default;
|
||||
};
|
||||
|
||||
template <>
|
||||
class output_buffer<subtractive_share<gf2, 0>> : public packed_share_output<gf2, 0>
|
||||
{
|
||||
using base = packed_share_output<gf2, 0>;
|
||||
public:
|
||||
using size_type = typename base::size_type;
|
||||
explicit output_buffer(size_type size) : base(size) { }
|
||||
output_buffer(output_buffer &&) noexcept = default;
|
||||
output_buffer(const output_buffer &) = delete;
|
||||
output_buffer & operator=(output_buffer &&) noexcept = default;
|
||||
output_buffer & operator=(const output_buffer &) = delete;
|
||||
~output_buffer() noexcept = default;
|
||||
};
|
||||
|
||||
template <>
|
||||
class output_buffer<subtractive_share<gf2, 1>> : public packed_share_output<gf2, 1>
|
||||
{
|
||||
using base = packed_share_output<gf2, 1>;
|
||||
public:
|
||||
using size_type = typename base::size_type;
|
||||
explicit output_buffer(size_type size) : base(size) { }
|
||||
output_buffer(output_buffer &&) noexcept = default;
|
||||
output_buffer(const output_buffer &) = delete;
|
||||
output_buffer & operator=(output_buffer &&) noexcept = default;
|
||||
output_buffer & operator=(const output_buffer &) = delete;
|
||||
~output_buffer() noexcept = default;
|
||||
};
|
||||
|
||||
template <>
|
||||
class output_buffer<subtractive_share<gf22, 0>> : public packed_share_output<gf22, 0>
|
||||
{
|
||||
using base = packed_share_output<gf22, 0>;
|
||||
public:
|
||||
using size_type = typename base::size_type;
|
||||
explicit output_buffer(size_type size) : base(size) { }
|
||||
output_buffer(output_buffer &&) noexcept = default;
|
||||
output_buffer(const output_buffer &) = delete;
|
||||
output_buffer & operator=(output_buffer &&) noexcept = default;
|
||||
output_buffer & operator=(const output_buffer &) = delete;
|
||||
~output_buffer() noexcept = default;
|
||||
};
|
||||
|
||||
template <>
|
||||
class output_buffer<subtractive_share<gf22, 1>> : public packed_share_output<gf22, 1>
|
||||
{
|
||||
using base = packed_share_output<gf22, 1>;
|
||||
public:
|
||||
using size_type = typename base::size_type;
|
||||
explicit output_buffer(size_type size) : base(size) { }
|
||||
output_buffer(output_buffer &&) noexcept = default;
|
||||
output_buffer(const output_buffer &) = delete;
|
||||
output_buffer & operator=(output_buffer &&) noexcept = default;
|
||||
output_buffer & operator=(const output_buffer &) = delete;
|
||||
~output_buffer() noexcept = default;
|
||||
};
|
||||
|
||||
template <>
|
||||
class output_buffer<subtractive_share<gf24, 0>> : public packed_share_output<gf24, 0>
|
||||
{
|
||||
using base = packed_share_output<gf24, 0>;
|
||||
public:
|
||||
using size_type = typename base::size_type;
|
||||
explicit output_buffer(size_type size) : base(size) { }
|
||||
output_buffer(output_buffer &&) noexcept = default;
|
||||
output_buffer(const output_buffer &) = delete;
|
||||
output_buffer & operator=(output_buffer &&) noexcept = default;
|
||||
output_buffer & operator=(const output_buffer &) = delete;
|
||||
~output_buffer() noexcept = default;
|
||||
};
|
||||
|
||||
template <>
|
||||
class output_buffer<subtractive_share<gf24, 1>> : public packed_share_output<gf24, 1>
|
||||
{
|
||||
using base = packed_share_output<gf24, 1>;
|
||||
public:
|
||||
using size_type = typename base::size_type;
|
||||
explicit output_buffer(size_type size) : base(size) { }
|
||||
output_buffer(output_buffer &&) noexcept = default;
|
||||
output_buffer(const output_buffer &) = delete;
|
||||
output_buffer & operator=(output_buffer &&) noexcept = default;
|
||||
output_buffer & operator=(const output_buffer &) = delete;
|
||||
~output_buffer() noexcept = default;
|
||||
};
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#include "dpf/shamir.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace detail
|
||||
{
|
||||
|
||||
/// @brief Shamir inverse for `gf2n`. Points `1 .. N` must stay distinct, so
|
||||
/// `N < 2^Bits`.
|
||||
template <unsigned Bits>
|
||||
struct shamir_field<gf2n<Bits>> : std::true_type
|
||||
{
|
||||
/// @brief `a^{-1}` by Fermat, `a^{2^Bits - 2}`.
|
||||
/// @param a a nonzero field element
|
||||
/// @return `a^{-1}`
|
||||
/// @throws std::invalid_argument if `a` is zero
|
||||
static gf2n<Bits> inv(gf2n<Bits> a)
|
||||
{
|
||||
return gf2n<Bits>{gf2_detail::inv<Bits>(a.raw())};
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace detail
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_GF2_HPP__
|
||||
277
include/dpf/gilboa.hpp
Normal file
277
include/dpf/gilboa.hpp
Normal file
|
|
@ -0,0 +1,277 @@
|
|||
/// @file dpf/gilboa.hpp
|
||||
/// @brief One-off Gilboa multiplication from correlated bit×ring triples.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_GILBOA_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_GILBOA_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <stdexcept>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/beaver.hpp"
|
||||
#include "dpf/edabit.hpp"
|
||||
#include "dpf/ot_pack.hpp"
|
||||
#include "dpf/random.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace gilboa
|
||||
{
|
||||
|
||||
/// @brief Clear product of additive shares (oracle only).
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
Ring mul_clear(Ring x0, Ring x1, Ring y0, Ring y1)
|
||||
{
|
||||
return static_cast<Ring>((x0 + x1) * (y0 + y1));
|
||||
}
|
||||
|
||||
template <typename Ring>
|
||||
struct product_shares
|
||||
{
|
||||
Ring z0{};
|
||||
Ring z1{};
|
||||
};
|
||||
|
||||
/// @brief Oracle: reconstruct, multiply, re-share. Not a party protocol.
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
product_shares<Ring> mul_dealer(Ring x0, Ring x1, Ring y0, Ring y1)
|
||||
{
|
||||
const Ring z = mul_clear(x0, x1, y0, y1);
|
||||
const Ring m = dpf::uniform_sample<Ring>();
|
||||
return product_shares<Ring>{m, static_cast<Ring>(z - m)};
|
||||
}
|
||||
|
||||
/// @brief Per-bit messages one party produces before the peer exchange.
|
||||
/// @details Factor `x` must be held as bits (party 1 share 0, or XOR bit
|
||||
/// shares). Additive `x` shares are bit-sliced locally — correct when
|
||||
/// one party holds the clear factor.
|
||||
template <typename Ring>
|
||||
struct gilboa_round
|
||||
{
|
||||
std::vector<Ring> d_share; ///< additive share of x_bit - a
|
||||
std::vector<Ring> e_share; ///< additive share of y - b
|
||||
std::vector<ot::bit_ring_triple<Ring>> triples;
|
||||
Ring x_share{};
|
||||
Ring y_share{};
|
||||
unsigned party = 0;
|
||||
};
|
||||
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
gilboa_round<Ring> mul_from_ot_begin(ot::pack & pack, Ring x_share, Ring y_share,
|
||||
unsigned party, unsigned bits = 64)
|
||||
{
|
||||
if (bits == 0 || bits > 8u * sizeof(Ring))
|
||||
throw std::invalid_argument("gilboa bits");
|
||||
if (pack.remaining_bit_ring() < bits)
|
||||
throw std::runtime_error("gilboa::mul_from_ot: need bit×ring triples");
|
||||
gilboa_round<Ring> r;
|
||||
r.x_share = x_share;
|
||||
r.y_share = y_share;
|
||||
r.party = party;
|
||||
r.d_share.resize(bits);
|
||||
r.e_share.resize(bits);
|
||||
r.triples.reserve(bits);
|
||||
for (unsigned i = 0; i < bits; ++i)
|
||||
{
|
||||
auto t = pack.take_bit_ring<Ring>();
|
||||
const Ring x_bit = static_cast<Ring>(
|
||||
(static_cast<std::uint64_t>(x_share) >> i) & 1u);
|
||||
r.d_share[i] = static_cast<Ring>(x_bit - Ring{t.a});
|
||||
r.e_share[i] = static_cast<Ring>(y_share - t.b);
|
||||
r.triples.push_back(t);
|
||||
}
|
||||
return r;
|
||||
}
|
||||
|
||||
/// @brief Finish after exchanging additive `d` and `e` with the peer.
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
Ring mul_from_ot_finish(const gilboa_round<Ring> & r,
|
||||
const std::vector<Ring> & peer_d, const std::vector<Ring> & peer_e)
|
||||
{
|
||||
if (peer_d.size() != r.d_share.size() || peer_e.size() != r.e_share.size())
|
||||
throw std::invalid_argument("gilboa finish size");
|
||||
Ring acc{};
|
||||
for (std::size_t i = 0; i < r.d_share.size(); ++i)
|
||||
{
|
||||
const auto & t = r.triples[i];
|
||||
const Ring d = static_cast<Ring>(r.d_share[i] + peer_d[i]);
|
||||
const Ring e = static_cast<Ring>(r.e_share[i] + peer_e[i]);
|
||||
Ring z = t.c;
|
||||
z = static_cast<Ring>(z + d * t.b);
|
||||
z = static_cast<Ring>(z + e * Ring{t.a});
|
||||
if (r.party == 0)
|
||||
z = static_cast<Ring>(z + d * e);
|
||||
acc = static_cast<Ring>(acc + (z << i));
|
||||
}
|
||||
return acc;
|
||||
}
|
||||
|
||||
/// @brief One bit×ring product of an additive 0/1 value with `y`.
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
product_shares<Ring> mul_bit_from_ot(ot::pack & pack0, ot::pack & pack1,
|
||||
Ring bit0, Ring bit1, Ring y0, Ring y1)
|
||||
{
|
||||
auto t0 = pack0.template take_bit_ring<Ring>();
|
||||
auto t1 = pack1.template take_bit_ring<Ring>();
|
||||
const Ring d0 = static_cast<Ring>(bit0 - Ring{t0.a});
|
||||
const Ring d1 = static_cast<Ring>(bit1 - Ring{t1.a});
|
||||
const Ring e0 = static_cast<Ring>(y0 - t0.b);
|
||||
const Ring e1 = static_cast<Ring>(y1 - t1.b);
|
||||
const Ring d = static_cast<Ring>(d0 + d1);
|
||||
const Ring e = static_cast<Ring>(e0 + e1);
|
||||
auto acc = [&](const ot::bit_ring_triple<Ring> & t, unsigned party) {
|
||||
Ring z = t.c;
|
||||
z = static_cast<Ring>(z + d * t.b);
|
||||
z = static_cast<Ring>(z + e * Ring{t.a});
|
||||
if (party == 0)
|
||||
z = static_cast<Ring>(z + d * e);
|
||||
return z;
|
||||
};
|
||||
return product_shares<Ring>{acc(t0, 0), acc(t1, 1)};
|
||||
}
|
||||
|
||||
/// @brief General additive `x`: A2B, daBit B2A per bit, then bit×ring Gilboa.
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
product_shares<Ring> mul_from_ot_pair(ot::pack & pack0, ot::pack & pack1,
|
||||
Ring x0, Ring x1, Ring y0, Ring y1, unsigned bits = 64)
|
||||
{
|
||||
if (pack0.remaining_b2a() < bits || pack1.remaining_b2a() < bits
|
||||
|| pack0.remaining_bit_ring() < bits || pack1.remaining_bit_ring() < bits)
|
||||
throw std::runtime_error("gilboa::mul_from_ot_pair: need dabits and bit×ring");
|
||||
auto eda = edabit::sample_edabit_pair<Ring>(bits);
|
||||
auto bits_xy = edabit::a2b_gmw_pair(eda, x0, x1);
|
||||
const auto & bx0 = bits_xy.first;
|
||||
const auto & bx1 = bits_xy.second;
|
||||
Ring z0{}, z1{};
|
||||
for (unsigned i = 0; i < bits; ++i)
|
||||
{
|
||||
auto d0 = pack0.template take_dabit<Ring>();
|
||||
auto d1 = pack1.template take_dabit<Ring>();
|
||||
const std::uint8_t b0 = edabit::detail::get_bit(bx0, i);
|
||||
const std::uint8_t b1 = edabit::detail::get_bit(bx1, i);
|
||||
const std::uint8_t mask = static_cast<std::uint8_t>(
|
||||
(b0 ^ d0.bit) ^ (b1 ^ d1.bit));
|
||||
const Ring a0 = edabit::b2a_party_bit(b0, d0, mask, 0, 0);
|
||||
const Ring a1 = edabit::b2a_party_bit(b1, d1, mask, 1, 0);
|
||||
auto part = mul_bit_from_ot(pack0, pack1, a0, a1, y0, y1);
|
||||
z0 = static_cast<Ring>(z0 + (part.z0 << i));
|
||||
z1 = static_cast<Ring>(z1 + (part.z1 << i));
|
||||
}
|
||||
return product_shares<Ring>{z0, z1};
|
||||
}
|
||||
|
||||
/// @brief Convenience: one party after peer messages are known (same as finish).
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
Ring mul_from_ot(ot::pack & pack, Ring x_share, Ring y_share, unsigned party,
|
||||
const std::vector<Ring> & peer_d, const std::vector<Ring> & peer_e,
|
||||
unsigned bits = 64)
|
||||
{
|
||||
auto r = mul_from_ot_begin(pack, x_share, y_share, party, bits);
|
||||
return mul_from_ot_finish(r, peer_d, peer_e);
|
||||
}
|
||||
|
||||
/// @brief Honest-dealer tape: `session::sample` then `export_party`.
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
beavers::party_tape<Ring> fill_tape_dealer(beavers::session<Ring> & s,
|
||||
unsigned party)
|
||||
{
|
||||
s.sample();
|
||||
return s.export_party(party);
|
||||
}
|
||||
|
||||
/// @brief Sample λ from correlated bit×ring `b` shares; monomials are Π λ^e.
|
||||
/// @details One λ per wire (repeated factors share it). Both packs are consumed
|
||||
/// in lockstep. `sample` is driven by those triples, not overwritten after.
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::pair<beavers::party_tape<Ring>, beavers::party_tape<Ring>>
|
||||
fill_tape_ot_pair(beavers::session<Ring> & s, ot::pack & pack0, ot::pack & pack1)
|
||||
{
|
||||
std::size_t calls = 0;
|
||||
s.sample([&]() {
|
||||
++calls;
|
||||
return Ring{1};
|
||||
});
|
||||
std::size_t nblind = 0;
|
||||
{
|
||||
const auto probe = s.export_party(0);
|
||||
for (auto ready : probe.lambda_ready)
|
||||
nblind += ready ? 1u : 0u;
|
||||
}
|
||||
s.clear_sample();
|
||||
struct sampler
|
||||
{
|
||||
ot::pack * p0;
|
||||
ot::pack * p1;
|
||||
std::size_t nblind;
|
||||
std::size_t call = 0;
|
||||
Ring saved{};
|
||||
Ring operator()()
|
||||
{
|
||||
if (p0->remaining_bit_ring() == 0 || p1->remaining_bit_ring() == 0)
|
||||
throw std::runtime_error("fill_tape_ot: need bit×ring triples");
|
||||
if (call < nblind * 2)
|
||||
{
|
||||
if ((call & 1u) == 0)
|
||||
{
|
||||
auto t0 = p0->template take_bit_ring<Ring>();
|
||||
auto t1 = p1->template take_bit_ring<Ring>();
|
||||
saved = t0.b;
|
||||
++call;
|
||||
return static_cast<Ring>(t0.b + t1.b);
|
||||
}
|
||||
++call;
|
||||
return saved;
|
||||
}
|
||||
auto t0 = p0->template take_bit_ring<Ring>();
|
||||
auto t1 = p1->template take_bit_ring<Ring>();
|
||||
++call;
|
||||
(void)t1;
|
||||
return t0.c;
|
||||
}
|
||||
};
|
||||
s.sample(sampler{&pack0, &pack1, nblind});
|
||||
(void)calls;
|
||||
auto t0 = s.export_party(0);
|
||||
auto t1 = s.export_party(1);
|
||||
pack0.stash_party_tape(t0);
|
||||
pack1.stash_party_tape(t1);
|
||||
return {t0, t1};
|
||||
}
|
||||
|
||||
/// @brief Read the party view `fill_tape_ot_pair` stashed in `pack`.
|
||||
/// @details Does not sample. A pack that was not produced by the pair sampler
|
||||
/// has no view to read.
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
beavers::party_tape<Ring> fill_tape_ot(beavers::session<Ring> & s, ot::pack & pack,
|
||||
unsigned party)
|
||||
{
|
||||
if (party > 1)
|
||||
throw std::invalid_argument("fill_tape_ot party");
|
||||
if (!pack.has_party_tape())
|
||||
throw std::runtime_error(
|
||||
"fill_tape_ot: pack has no party view; call fill_tape_ot_pair");
|
||||
auto tape = pack.template load_party_tape<Ring>();
|
||||
if (tape.lambda.size() != s.wire_count())
|
||||
throw std::invalid_argument("fill_tape_ot: tape does not match session");
|
||||
(void)party;
|
||||
return tape;
|
||||
}
|
||||
|
||||
} // namespace gilboa
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_GILBOA_HPP__
|
||||
1143
include/dpf/grow.hpp
Normal file
1143
include/dpf/grow.hpp
Normal file
File diff suppressed because it is too large
Load diff
132
include/dpf/grow_ds.hpp
Normal file
132
include/dpf/grow_ds.hpp
Normal file
|
|
@ -0,0 +1,132 @@
|
|||
/// @file dpf/grow_ds.hpp
|
||||
/// @brief Doerner–Shelat interactive / joint grow: `extend_ds` and
|
||||
/// `add_output_ds` on path-memoizer frontiers.
|
||||
/// @details Joint (2+1 / local) view holds both keys and both memoizers. One
|
||||
/// correction-word round uses `ds_advance_level`; leaf planting reuses
|
||||
/// dealer `grow_impl` with the opened CW. A socket backend swaps in a
|
||||
/// `CwProtocol` that exchanges blinds without revealing the peer seed.
|
||||
/// @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_DPF_GROW_DS_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_GROW_DS_HPP__
|
||||
|
||||
#include <array>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/doerner_shelat.hpp"
|
||||
#include "dpf/grow.hpp"
|
||||
#include "dpf/path_memoizer.hpp"
|
||||
#include "dpf/utils.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
/// @brief One interactive interior level from memoizer frontiers, then plant
|
||||
/// `specs` via the same assembly as dealer `extend`.
|
||||
/// @details `x0` / `x1` are XOR shares of the programmed point (use `(x, 0)` in
|
||||
/// a local joint test). Memoizers must already be filled for the clear
|
||||
/// point `x0 ⊕ x1` through the old depth.
|
||||
/// \complexity One `ds_advance_level` (two PRG expands + `prepare_level` /
|
||||
/// `open_cw` / AND opens) plus the same leaf plants as dealer
|
||||
/// `extend`. No O(d) rewalk when memoizers are warm.
|
||||
/// \rounds One interactive CW round for the new level (local_cw_protocol opens
|
||||
/// in-process; a socket `CwProtocol` is one peer exchange round).
|
||||
/// \communication Local: none on the wire. Networked: one level's blinds, CW
|
||||
/// shares, and advice (same shape as one `point_party` level),
|
||||
/// plus leaf pads when planting.
|
||||
template <typename K0, typename K1, typename Memo0, typename Memo1,
|
||||
typename InputT, typename CwProtocol, typename... Specs>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto extend_ds(const K0 & k0, const K1 & k1, Memo0 & m0, Memo1 & m1, InputT x0,
|
||||
InputT x1, CwProtocol & proto, Specs &&... specs)
|
||||
{
|
||||
using old_key = detail::grow_impl::bare_key_t<K0>;
|
||||
static_assert(std::is_same_v<old_key, detail::grow_impl::bare_key_t<K1>>,
|
||||
"extend_ds: both keys must have the same type");
|
||||
using input_type = typename old_key::input_type;
|
||||
using node = typename old_key::interior_node;
|
||||
using interior = typename old_key::interior_prg;
|
||||
|
||||
input_type xx0 = static_cast<input_type>(x0);
|
||||
input_type xx1 = static_cast<input_type>(x1);
|
||||
utils::flip_msb_if_signed_integral(xx0);
|
||||
// Party 1 share is not MSB-flipped in DS (same as make_dpf_doerner_shelat).
|
||||
const input_type x = utils::xor_input_shares(xx0, xx1);
|
||||
|
||||
const old_key & bk0 = static_cast<const old_key &>(k0);
|
||||
const old_key & bk1 = static_cast<const old_key &>(k1);
|
||||
|
||||
const bool bit = detail::grow_impl::bit_at(x, old_key::depth);
|
||||
|
||||
node s0{};
|
||||
node s1{};
|
||||
std::array<bool, old_key::depth == 0 ? 1 : old_key::depth> path{};
|
||||
detail::grow_impl::frontier_from_memos(bk0, bk1, m0, m1, x, old_key::depth, s0,
|
||||
s1, old_key::depth == 0 ? nullptr : path.data());
|
||||
|
||||
detail::ds_gen_state<node> st;
|
||||
st.init(s0, s1);
|
||||
|
||||
const std::size_t level = old_key::depth;
|
||||
const std::size_t new_depth = old_key::depth + 1;
|
||||
auto mask = old_key::msb_mask;
|
||||
for (std::size_t i = 0; i < level; ++i)
|
||||
mask >>= 1;
|
||||
|
||||
node cw{};
|
||||
psnip_uint8_t advice = 0;
|
||||
detail::ds_advance_level<interior>(st, xx0, xx1, mask, level, new_depth,
|
||||
proto, cw, advice);
|
||||
|
||||
node ns0 = st.seed0();
|
||||
node ns1 = st.seed1();
|
||||
|
||||
return detail::grow_impl::grow_impl<true, true>(bk0, bk1, &m0, &m1, bit, x,
|
||||
true, &cw, &advice, &ns0, &ns1, std::forward<Specs>(specs)...);
|
||||
}
|
||||
|
||||
/// @brief Plant outputs on existing levels using memoizer seeds. Joint leaf
|
||||
/// construction matches dealer `add_output`; `proto.open_leaf_group` is
|
||||
/// invoked so a non-local protocol can hide the clear point.
|
||||
/// \complexity O(1) frontier reads plus leaf plants. No new interior CW.
|
||||
/// \rounds Leaf-open only (local: one `open_leaf_group` callback; networked:
|
||||
/// the leaf pad / mux pattern of `point_party`).
|
||||
/// \communication none for `local_cw_protocol`; otherwise leaf pads and the
|
||||
/// leaf CW open.
|
||||
template <typename K0, typename K1, typename Memo0, typename Memo1,
|
||||
typename InputT, typename CwProtocol, typename... Specs>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto add_output_ds(const K0 & k0, const K1 & k1, Memo0 & m0, Memo1 & m1,
|
||||
InputT x0, InputT x1, CwProtocol & proto, Specs &&... specs)
|
||||
{
|
||||
using old_key = detail::grow_impl::bare_key_t<K0>;
|
||||
using input_type = typename old_key::input_type;
|
||||
const old_key & bk0 = static_cast<const old_key &>(k0);
|
||||
const old_key & bk1 = static_cast<const old_key &>(k1);
|
||||
|
||||
input_type xx0 = static_cast<input_type>(x0);
|
||||
input_type xx1 = static_cast<input_type>(x1);
|
||||
utils::flip_msb_if_signed_integral(xx0);
|
||||
input_type x{};
|
||||
proto.open_leaf_group(xx0, xx1, [&](input_type sx0, input_type sx1) {
|
||||
x = utils::xor_input_shares(sx0, sx1);
|
||||
});
|
||||
|
||||
return detail::grow_impl::grow_impl<false, true>(bk0, bk1, &m0, &m1,
|
||||
/*bit=*/false, x, false,
|
||||
static_cast<typename old_key::interior_node *>(nullptr),
|
||||
static_cast<psnip_uint8_t *>(nullptr),
|
||||
static_cast<typename old_key::interior_node *>(nullptr),
|
||||
static_cast<typename old_key::interior_node *>(nullptr),
|
||||
std::forward<Specs>(specs)...);
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_GROW_DS_HPP__
|
||||
202
include/dpf/idpf_agg.hpp
Normal file
202
include/dpf/idpf_agg.hpp
Normal file
|
|
@ -0,0 +1,202 @@
|
|||
/// @file dpf/idpf_agg.hpp
|
||||
/// @brief Max and k-th order statistic on a list of incremental DPF keys.
|
||||
/// @details Each secret value is one `idpf` with a unit payload on every
|
||||
/// prefix length. Servers add prefix shares along the live spine
|
||||
/// with `eval_until`. Communication tracks the bit length of the
|
||||
/// domain, not how many secret inputs were summed
|
||||
/// (Cheng–Mitrokotsa–Zhang–Hartmann, ePrint 2024/1190, on the
|
||||
/// S&P 2021 incremental DPF / Poplar walk).
|
||||
/// @see dpf/eval_until.hpp, dpf/placement.hpp (`idpf`), examples/applications/idpf_agg.cpp
|
||||
/// @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_DPF_IDPF_AGG_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_IDPF_AGG_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <stdexcept>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/eval_until.hpp"
|
||||
#include "dpf/placement.hpp"
|
||||
#include "dpf/secret_share.hpp"
|
||||
#include "dpf/utils.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
/// @brief Unit payload on every prefix length `1 .. N`.
|
||||
template <std::size_t N, typename Beta = std::uint64_t, std::size_t... Is>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr auto idpf_ones_impl(std::index_sequence<Is...>)
|
||||
{
|
||||
return idpf(((void)Is, Beta{1})...);
|
||||
}
|
||||
|
||||
/// @brief `idpf(1, 1, …, 1)` of length `N` (one unit per prefix length).
|
||||
template <std::size_t N, typename Beta = std::uint64_t>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr auto idpf_ones()
|
||||
{
|
||||
static_assert(N > 0, "idpf_ones: N must be positive");
|
||||
return idpf_ones_impl<N, Beta>(std::make_index_sequence<N>{});
|
||||
}
|
||||
|
||||
namespace detail
|
||||
{
|
||||
namespace idpf_agg_detail
|
||||
{
|
||||
|
||||
template <typename Share, typename = void>
|
||||
struct has_raw_member : std::false_type {};
|
||||
template <typename Share>
|
||||
struct has_raw_member<Share,
|
||||
std::void_t<decltype(std::declval<const Share &>().raw())>>
|
||||
: std::true_type {};
|
||||
|
||||
template <typename Share>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
auto share_raw(const Share & s)
|
||||
{
|
||||
if constexpr (has_raw_member<Share>::value)
|
||||
return s.raw();
|
||||
else
|
||||
return s;
|
||||
}
|
||||
|
||||
template <typename Key0T, typename Key1T>
|
||||
auto open_prefix_counts(std::vector<idpf_eval_ctx<Key0T>> & ctx0,
|
||||
std::vector<idpf_eval_ctx<Key1T>> & ctx1, std::size_t level,
|
||||
const std::vector<typename unwrap_party_key_t<Key0T>::input_type> & prefs)
|
||||
{
|
||||
std::vector<std::uint64_t> counts(prefs.size(), 0);
|
||||
for (std::size_t i = 0; i < ctx0.size(); ++i)
|
||||
{
|
||||
auto s0 = eval_until(ctx0[i], level, prefs);
|
||||
auto s1 = eval_until(ctx1[i], level, prefs);
|
||||
for (std::size_t j = 0; j < prefs.size(); ++j)
|
||||
{
|
||||
const auto opened = reconstruct(s0[j], s1[j]);
|
||||
counts[j] += static_cast<std::uint64_t>(share_raw(opened));
|
||||
}
|
||||
}
|
||||
return counts;
|
||||
}
|
||||
|
||||
template <typename Key0T, typename Key1T, typename InputT>
|
||||
void retain_all(std::vector<idpf_eval_ctx<Key0T>> & ctx0,
|
||||
std::vector<idpf_eval_ctx<Key1T>> & ctx1, InputT prefix)
|
||||
{
|
||||
for (std::size_t i = 0; i < ctx0.size(); ++i)
|
||||
{
|
||||
ctx0[i].retain(prefix);
|
||||
ctx1[i].retain(prefix);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace idpf_agg_detail
|
||||
} // namespace detail
|
||||
|
||||
/// @brief Walk both parties' idpf contexts to the numeric maximum.
|
||||
/// @details At each bit, keep the `1`-child when it holds any mass; otherwise
|
||||
/// keep the `0`-child. That is the MSB-first spine of the max value.
|
||||
/// \complexity O(n · N · depth) party work for `n` keys on an `N`-bit domain:
|
||||
/// two `eval_until` calls per key per level (both children), then
|
||||
/// `retain`. Communication of opened counts is O(depth), independent
|
||||
/// of `n` (ePrint 2024/1190).
|
||||
template <typename Key0T, typename Key1T>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto idpf_agg_max(const std::vector<Key0T> & keys0,
|
||||
const std::vector<Key1T> & keys1)
|
||||
{
|
||||
using key_type = unwrap_party_key_t<Key0T>;
|
||||
using input_type = typename key_type::input_type;
|
||||
constexpr auto bits = key_type::input_bits;
|
||||
static_assert(std::is_same_v<key_type, unwrap_party_key_t<Key1T>>,
|
||||
"idpf_agg_max: party keys must unwrap to the same DPF type");
|
||||
if (keys0.size() != keys1.size() || keys0.empty())
|
||||
throw std::invalid_argument("idpf_agg_max: nonempty equal key lists");
|
||||
|
||||
std::vector<idpf_eval_ctx<Key0T>> ctx0;
|
||||
std::vector<idpf_eval_ctx<Key1T>> ctx1;
|
||||
ctx0.reserve(keys0.size());
|
||||
ctx1.reserve(keys1.size());
|
||||
for (std::size_t i = 0; i < keys0.size(); ++i)
|
||||
{
|
||||
ctx0.emplace_back(keys0[i]);
|
||||
ctx1.emplace_back(keys1[i]);
|
||||
}
|
||||
|
||||
input_type prefix{0};
|
||||
for (std::size_t level = 1; level <= bits; ++level)
|
||||
{
|
||||
const input_type left = static_cast<input_type>(prefix << 1);
|
||||
const input_type right = static_cast<input_type>((prefix << 1) | 1);
|
||||
std::vector<input_type> prefs{left, right};
|
||||
auto counts = detail::idpf_agg_detail::open_prefix_counts(ctx0, ctx1,
|
||||
level, prefs);
|
||||
prefix = (counts[1] > 0) ? right : left;
|
||||
detail::idpf_agg_detail::retain_all(ctx0, ctx1, prefix);
|
||||
}
|
||||
return prefix;
|
||||
}
|
||||
|
||||
/// @brief Walk both parties' idpf contexts to the k-th largest (1-based).
|
||||
/// @details `k == 1` is the maximum; `k == n` is the minimum. At each bit,
|
||||
/// take the `1`-child when its count is at least `k`; otherwise
|
||||
/// subtract that count and take the `0`-child.
|
||||
/// \complexity Same as `idpf_agg_max`: O(n · N · depth) local work,
|
||||
/// O(depth) opened counts.
|
||||
template <typename Key0T, typename Key1T>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto idpf_agg_kth(const std::vector<Key0T> & keys0,
|
||||
const std::vector<Key1T> & keys1, std::size_t k)
|
||||
{
|
||||
using key_type = unwrap_party_key_t<Key0T>;
|
||||
using input_type = typename key_type::input_type;
|
||||
constexpr auto bits = key_type::input_bits;
|
||||
static_assert(std::is_same_v<key_type, unwrap_party_key_t<Key1T>>,
|
||||
"idpf_agg_kth: party keys must unwrap to the same DPF type");
|
||||
if (keys0.size() != keys1.size() || keys0.empty())
|
||||
throw std::invalid_argument("idpf_agg_kth: nonempty equal key lists");
|
||||
if (k == 0 || k > keys0.size())
|
||||
throw std::invalid_argument("idpf_agg_kth: k out of range");
|
||||
|
||||
std::vector<idpf_eval_ctx<Key0T>> ctx0;
|
||||
std::vector<idpf_eval_ctx<Key1T>> ctx1;
|
||||
ctx0.reserve(keys0.size());
|
||||
ctx1.reserve(keys1.size());
|
||||
for (std::size_t i = 0; i < keys0.size(); ++i)
|
||||
{
|
||||
ctx0.emplace_back(keys0[i]);
|
||||
ctx1.emplace_back(keys1[i]);
|
||||
}
|
||||
|
||||
input_type prefix{0};
|
||||
for (std::size_t level = 1; level <= bits; ++level)
|
||||
{
|
||||
const input_type left = static_cast<input_type>(prefix << 1);
|
||||
const input_type right = static_cast<input_type>((prefix << 1) | 1);
|
||||
std::vector<input_type> prefs{left, right};
|
||||
auto counts = detail::idpf_agg_detail::open_prefix_counts(ctx0, ctx1,
|
||||
level, prefs);
|
||||
if (counts[1] >= k)
|
||||
prefix = right;
|
||||
else
|
||||
{
|
||||
k -= counts[1];
|
||||
prefix = left;
|
||||
}
|
||||
detail::idpf_agg_detail::retain_all(ctx0, ctx1, prefix);
|
||||
}
|
||||
return prefix;
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_IDPF_AGG_HPP__
|
||||
700
include/dpf/iknp.hpp
Normal file
700
include/dpf/iknp.hpp
Normal file
|
|
@ -0,0 +1,700 @@
|
|||
/// @file dpf/iknp.hpp
|
||||
/// @brief Semi-honest IKNP OT extension and the pads a DPF dealer would sample.
|
||||
/// @details Base OTs are Chou–Orlandi (LATINCRYPT 2015, ePrint 2015/267) on
|
||||
/// P-256. Extension follows Ishai, Kilian, Nissim, and Petrank,
|
||||
/// CRYPTO 2003, with fixed-key AES as the correlation-robust hash.
|
||||
/// `sample` returns this party's shares of random bit triples,
|
||||
/// bit×block triples, comparison B2A pads, and Doerner–Shelat
|
||||
/// correction-word pads.
|
||||
///
|
||||
/// Ideal functionality of `sample` (semi-honest, two parties):
|
||||
/// - **Inputs.** Both parties pass the same lengths `(nblock, nbit, nb2a, ncw)`
|
||||
/// and call on the peer link before any other walk traffic.
|
||||
/// - **Outputs.** Party `i` receives shares such that
|
||||
/// - blocks: `(a0⊕a1)·(b0⊕b1) = c0⊕c1` (bit × 128-bit block);
|
||||
/// - bits: `(a0⊕a1)∧(b0⊕b1) = c0⊕c1`;
|
||||
/// - b2a: `(add0+add1) mod 2^64 = r0⊕r1` (0 or 1);
|
||||
/// - cw: `gamma0⊕gamma1 = (bit1·rand0)⊕(bit0·rand1)` (XOR shares — neither
|
||||
/// party learns the peer pad bit; see `ds_sample_cw`).
|
||||
/// - **Hidden.** Peer seeds, peer choice bits used only as OT choice, and the
|
||||
/// peer's pad bit (so an opened blind `path⊕pad` does not open the path).
|
||||
///
|
||||
/// Cost of `sample` with tape `T = nblock + nbit + ncw` and security
|
||||
/// parameter `κ = 128`: two Chou–Orlandi base sessions of `κ` OTs (P-256
|
||||
/// points, `Θ(κ)` scalar muls), then two OT-extension directions each
|
||||
/// sending `κ ⌈T/8⌉` bytes of U and `T · 16` bytes of correction, plus
|
||||
/// `ncw · 16` bytes of `gamma` and an optional B2A extension of length
|
||||
/// `nb2a`. See [tour_iknp](@ref tour_iknp) for the comparison with a p2
|
||||
/// dealer tape and with Half-Tree §5.2.
|
||||
|
||||
#ifndef LIBDPF_INCLUDE_DPF_IKNP_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_IKNP_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/net/channel.hpp"
|
||||
#include "dpf/p256.hpp"
|
||||
#include "dpf/prg_aes.hpp"
|
||||
#include "dpf/random.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace iknp
|
||||
{
|
||||
|
||||
struct block_share
|
||||
{
|
||||
std::uint8_t a = 0;
|
||||
simde__m128i b{};
|
||||
simde__m128i c{};
|
||||
};
|
||||
|
||||
struct bit_share
|
||||
{
|
||||
std::uint8_t a = 0;
|
||||
std::uint8_t b = 0;
|
||||
std::uint8_t c = 0;
|
||||
};
|
||||
|
||||
struct b2a_share
|
||||
{
|
||||
std::uint8_t r = 0;
|
||||
std::uint64_t add = 0;
|
||||
};
|
||||
|
||||
struct cw_share
|
||||
{
|
||||
simde__m128i rand{};
|
||||
simde__m128i gamma{};
|
||||
std::uint8_t bit = 0;
|
||||
};
|
||||
|
||||
struct material
|
||||
{
|
||||
std::vector<block_share> blocks;
|
||||
std::vector<bit_share> bits;
|
||||
std::vector<b2a_share> b2a;
|
||||
std::vector<cw_share> cws;
|
||||
};
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
inline constexpr std::size_t kappa = 128;
|
||||
inline constexpr std::size_t chunk_rows = 8192;
|
||||
|
||||
struct point_msg
|
||||
{
|
||||
std::uint8_t enc[33];
|
||||
};
|
||||
|
||||
struct role_state
|
||||
{
|
||||
bool ready = false;
|
||||
simde__m128i k0[kappa]{};
|
||||
simde__m128i k1[kappa]{};
|
||||
simde__m128i seed[kappa]{};
|
||||
std::uint8_t delta_bits[kappa]{};
|
||||
simde__m128i delta{};
|
||||
std::uint64_t rows = 0;
|
||||
};
|
||||
|
||||
inline simde__m128i xor_block(simde__m128i a, simde__m128i b)
|
||||
{
|
||||
return simde_mm_xor_si128(a, b);
|
||||
}
|
||||
|
||||
inline std::uint8_t lsb(simde__m128i x)
|
||||
{
|
||||
unsigned char b = 0;
|
||||
std::memcpy(&b, &x, 1);
|
||||
return static_cast<std::uint8_t>(b & 1u);
|
||||
}
|
||||
|
||||
inline simde__m128i bit_block(std::uint8_t bit)
|
||||
{
|
||||
simde__m128i z = simde_mm_setzero_si128();
|
||||
const unsigned char b = static_cast<unsigned char>(bit & 1u);
|
||||
std::memcpy(&z, &b, 1);
|
||||
return z;
|
||||
}
|
||||
|
||||
inline simde__m128i gate(std::uint8_t bit, simde__m128i block)
|
||||
{
|
||||
return (bit & 1u) ? block : simde_mm_setzero_si128();
|
||||
}
|
||||
|
||||
inline std::uint64_t low64(simde__m128i x)
|
||||
{
|
||||
std::uint64_t v = 0;
|
||||
std::memcpy(&v, &x, sizeof(v));
|
||||
return v;
|
||||
}
|
||||
|
||||
inline simde__m128i ot_hash(std::uint64_t index, simde__m128i row)
|
||||
{
|
||||
const prg::purpose_scope counted(prg::purpose::hash);
|
||||
const auto mixed = xor_block(row,
|
||||
simde_mm_set_epi64x(static_cast<std::int64_t>(index >> 32),
|
||||
static_cast<std::int64_t>(index)));
|
||||
return prg::aes128::eval(mixed,
|
||||
static_cast<psnip_uint32_t>(index * 0x9E3779B9u) ^ 0xA5A5u);
|
||||
}
|
||||
|
||||
inline void sample_scalar(std::uint64_t k[4])
|
||||
{
|
||||
for (;;)
|
||||
{
|
||||
for (int i = 0; i < 4; ++i)
|
||||
k[i] = dpf::uniform_sample<std::uint64_t>();
|
||||
const bool zero = (k[0] | k[1] | k[2] | k[3]) == 0;
|
||||
if (!zero && p256_detail::limbs_cmp(k, p256_detail::N, 4) < 0)
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
inline simde__m128i hash_point(const p256_detail::affine & p)
|
||||
{
|
||||
std::uint8_t enc[33]{};
|
||||
p256_detail::encode_point(enc, p);
|
||||
simde__m128i b0 = simde_mm_setzero_si128();
|
||||
simde__m128i b1 = simde_mm_setzero_si128();
|
||||
std::memcpy(&b0, enc, 16);
|
||||
std::memcpy(&b1, enc + 16, 16);
|
||||
const simde__m128i b2 = simde_mm_set_epi64x(0, enc[32]);
|
||||
const prg::purpose_scope counted(prg::purpose::hash);
|
||||
auto h = prg::aes128::eval(b0, 1);
|
||||
h = xor_block(h, prg::aes128::eval(b1, 2));
|
||||
return xor_block(h, prg::aes128::eval(b2, 3));
|
||||
}
|
||||
|
||||
inline simde__m128i pack_bits(const std::uint8_t * bits)
|
||||
{
|
||||
std::uint8_t packed[16]{};
|
||||
for (int i = 0; i < static_cast<int>(kappa); ++i)
|
||||
{
|
||||
if (bits[i] & 1u)
|
||||
packed[static_cast<unsigned>(i) >> 3] |=
|
||||
static_cast<std::uint8_t>(1u << (i & 7));
|
||||
}
|
||||
simde__m128i out = simde_mm_setzero_si128();
|
||||
std::memcpy(&out, packed, 16);
|
||||
return out;
|
||||
}
|
||||
|
||||
inline void expand_column(simde__m128i seed, std::uint64_t domain,
|
||||
std::uint8_t * dst, std::size_t nbytes)
|
||||
{
|
||||
const auto tweaked = xor_block(seed,
|
||||
simde_mm_set_epi64x(static_cast<std::int64_t>(domain), 0));
|
||||
const std::size_t nblocks = (nbytes + 15) / 16;
|
||||
std::vector<simde__m128i> buf(nblocks);
|
||||
if (nblocks > 0)
|
||||
prg::aes128::eval(tweaked, buf.data(),
|
||||
static_cast<psnip_uint32_t>(nblocks), 0);
|
||||
std::memcpy(dst, buf.data(), nbytes);
|
||||
}
|
||||
|
||||
/// @brief Transpose a `kappa × nrows` bit matrix packed by columns into rows.
|
||||
inline void transpose_rows(const std::uint8_t * cols, std::size_t nbytes,
|
||||
std::size_t nrows, simde__m128i * rows)
|
||||
{
|
||||
// Process 8 row-bits at a time when possible via byte gathers; fall back
|
||||
// per-row for the tail. Still O(kappa · nrows) but with tight inner loops.
|
||||
for (std::size_t j = 0; j < nrows; ++j)
|
||||
{
|
||||
std::uint8_t packed[16]{};
|
||||
const std::size_t byte = j >> 3;
|
||||
const auto mask = static_cast<std::uint8_t>(1u << (j & 7));
|
||||
for (int i = 0; i < static_cast<int>(kappa); ++i)
|
||||
{
|
||||
if (cols[static_cast<std::size_t>(i) * nbytes + byte] & mask)
|
||||
packed[static_cast<unsigned>(i) >> 3] |=
|
||||
static_cast<std::uint8_t>(1u << (i & 7));
|
||||
}
|
||||
std::memcpy(&rows[j], packed, 16);
|
||||
}
|
||||
}
|
||||
|
||||
inline void base_sender(net::channel & ch, simde__m128i k0[kappa],
|
||||
simde__m128i k1[kappa])
|
||||
{
|
||||
std::uint64_t a[4];
|
||||
sample_scalar(a);
|
||||
const auto A = p256_detail::point_scalarmul_limbs(
|
||||
p256_detail::generator_point(), a);
|
||||
point_msg am{};
|
||||
p256_detail::encode_point(am.enc, A);
|
||||
ch.send(net::msg::bytes, am);
|
||||
const auto bs = ch.recv_vec<point_msg>(net::msg::bytes);
|
||||
if (bs.size() != kappa)
|
||||
throw std::runtime_error("iknp base OT count");
|
||||
for (std::size_t i = 0; i < kappa; ++i)
|
||||
{
|
||||
const auto B = p256_detail::decode_strict(bs[i].enc, 33);
|
||||
k0[i] = hash_point(p256_detail::point_scalarmul_limbs(B, a));
|
||||
k1[i] = hash_point(p256_detail::point_scalarmul_limbs(
|
||||
p256_detail::point_sub(B, A), a));
|
||||
}
|
||||
}
|
||||
|
||||
inline void base_receiver(net::channel & ch, simde__m128i seed[kappa],
|
||||
std::uint8_t delta_bits[kappa], simde__m128i & delta)
|
||||
{
|
||||
const auto am = ch.recv<point_msg>(net::msg::bytes);
|
||||
const auto A = p256_detail::decode_strict(am.enc, 33);
|
||||
std::vector<point_msg> bs(kappa);
|
||||
for (std::size_t i = 0; i < kappa; ++i)
|
||||
{
|
||||
delta_bits[i] = static_cast<std::uint8_t>(
|
||||
dpf::uniform_sample<unsigned char>() & 1u);
|
||||
std::uint64_t r[4];
|
||||
sample_scalar(r);
|
||||
const auto R = p256_detail::point_scalarmul_limbs(
|
||||
p256_detail::generator_point(), r);
|
||||
const auto B = delta_bits[i]
|
||||
? p256_detail::point_add(A, R) : R;
|
||||
p256_detail::encode_point(bs[i].enc, B);
|
||||
seed[i] = hash_point(p256_detail::point_scalarmul_limbs(A, r));
|
||||
}
|
||||
delta = pack_bits(delta_bits);
|
||||
ch.send_vec(bs, net::msg::bytes);
|
||||
}
|
||||
|
||||
inline void ensure_sender(net::channel & ch, role_state & st)
|
||||
{
|
||||
if (st.ready)
|
||||
return;
|
||||
base_receiver(ch, st.seed, st.delta_bits, st.delta);
|
||||
st.ready = true;
|
||||
}
|
||||
|
||||
inline void ensure_receiver(net::channel & ch, role_state & st)
|
||||
{
|
||||
if (st.ready)
|
||||
return;
|
||||
base_sender(ch, st.k0, st.k1);
|
||||
st.ready = true;
|
||||
}
|
||||
|
||||
inline void extend_send(net::channel & ch, role_state & st, std::size_t n,
|
||||
std::vector<simde__m128i> & m0, std::vector<simde__m128i> & m1)
|
||||
{
|
||||
if (n == 0)
|
||||
return;
|
||||
ensure_sender(ch, st);
|
||||
m0.resize(n);
|
||||
m1.resize(n);
|
||||
std::size_t off = 0;
|
||||
while (off < n)
|
||||
{
|
||||
const std::size_t rows = std::min(chunk_rows, n - off);
|
||||
const std::size_t nbytes = (rows + 7) / 8;
|
||||
const auto u = ch.recv_bytes(net::msg::bytes);
|
||||
if (u.size() != kappa * nbytes)
|
||||
throw std::runtime_error("iknp extension length");
|
||||
std::vector<std::uint8_t> cols(kappa * nbytes);
|
||||
for (std::size_t i = 0; i < kappa; ++i)
|
||||
{
|
||||
expand_column(st.seed[i], st.rows, cols.data() + i * nbytes, nbytes);
|
||||
if (st.delta_bits[i])
|
||||
{
|
||||
for (std::size_t b = 0; b < nbytes; ++b)
|
||||
cols[i * nbytes + b] = static_cast<std::uint8_t>(
|
||||
cols[i * nbytes + b] ^ u[i * nbytes + b]);
|
||||
}
|
||||
}
|
||||
std::vector<simde__m128i> q(rows);
|
||||
transpose_rows(cols.data(), nbytes, rows, q.data());
|
||||
for (std::size_t j = 0; j < rows; ++j)
|
||||
{
|
||||
const auto index = st.rows + j;
|
||||
m0[off + j] = ot_hash(index, q[j]);
|
||||
m1[off + j] = ot_hash(index, xor_block(q[j], st.delta));
|
||||
}
|
||||
st.rows += rows;
|
||||
off += rows;
|
||||
}
|
||||
}
|
||||
|
||||
inline void extend_recv(net::channel & ch, role_state & st,
|
||||
const std::uint8_t * choices, std::size_t n,
|
||||
std::vector<simde__m128i> & masks)
|
||||
{
|
||||
if (n == 0)
|
||||
return;
|
||||
ensure_receiver(ch, st);
|
||||
masks.resize(n);
|
||||
std::size_t off = 0;
|
||||
while (off < n)
|
||||
{
|
||||
const std::size_t rows = std::min(chunk_rows, n - off);
|
||||
const std::size_t nbytes = (rows + 7) / 8;
|
||||
std::vector<std::uint8_t> xbytes(nbytes);
|
||||
for (std::size_t j = 0; j < rows; ++j)
|
||||
{
|
||||
if (choices[off + j] & 1u)
|
||||
xbytes[j >> 3] = static_cast<std::uint8_t>(
|
||||
xbytes[j >> 3] | (1u << (j & 7)));
|
||||
}
|
||||
std::vector<std::uint8_t> u(kappa * nbytes);
|
||||
std::vector<std::uint8_t> t0cols(kappa * nbytes);
|
||||
for (std::size_t i = 0; i < kappa; ++i)
|
||||
{
|
||||
std::vector<std::uint8_t> t1(nbytes);
|
||||
expand_column(st.k0[i], st.rows, t0cols.data() + i * nbytes, nbytes);
|
||||
expand_column(st.k1[i], st.rows, t1.data(), nbytes);
|
||||
for (std::size_t b = 0; b < nbytes; ++b)
|
||||
u[i * nbytes + b] = static_cast<std::uint8_t>(
|
||||
t0cols[i * nbytes + b] ^ t1[b] ^ xbytes[b]);
|
||||
}
|
||||
ch.send_bytes(net::msg::bytes, u.data(), u.size());
|
||||
std::vector<simde__m128i> t(rows);
|
||||
transpose_rows(t0cols.data(), nbytes, rows, t.data());
|
||||
for (std::size_t j = 0; j < rows; ++j)
|
||||
masks[off + j] = ot_hash(st.rows + j, t[j]);
|
||||
st.rows += rows;
|
||||
off += rows;
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Correlated OT correction: sender holds (m0,m1), payload Δ; receiver
|
||||
/// with choice χ gets m_χ ⊕ (χ·Δ). Parties get additive XOR shares of
|
||||
/// χ·Δ by taking sender share = m0 and receiver share = got ⊕ m0 path.
|
||||
/// @details Sender transmits `d = m0⊕m1⊕payload`. Receiver returns
|
||||
/// `choice ? mask⊕d : mask`. Sender's share is `m0`; receiver's is
|
||||
/// the returned value. Then `sender⊕receiver = choice·payload` when
|
||||
/// the OT is correct (`mask = m_choice`).
|
||||
inline void correct_ot(net::channel & ch, int me, bool i_am_sender,
|
||||
const std::vector<simde__m128i> & m0,
|
||||
const std::vector<simde__m128i> & m1,
|
||||
const std::vector<simde__m128i> & payloads,
|
||||
const std::vector<simde__m128i> & masks,
|
||||
const std::vector<std::uint8_t> & choices,
|
||||
std::vector<simde__m128i> & share)
|
||||
{
|
||||
const std::size_t n = i_am_sender ? m0.size() : masks.size();
|
||||
share.resize(n);
|
||||
if (n == 0)
|
||||
return;
|
||||
if (i_am_sender)
|
||||
{
|
||||
std::vector<simde__m128i> d(n);
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
d[i] = xor_block(xor_block(m0[i], m1[i]), payloads[i]);
|
||||
// Fixed order: party 0 always sends first when it is the sender;
|
||||
// when party 1 is the sender it sends (party 0 receives).
|
||||
if (me == 0)
|
||||
ch.send_vec(d, net::msg::bytes);
|
||||
else
|
||||
ch.send_vec(d, net::msg::bytes);
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
share[i] = m0[i];
|
||||
return;
|
||||
}
|
||||
auto d = ch.recv_vec<simde__m128i>(net::msg::bytes);
|
||||
if (d.size() != n)
|
||||
throw std::runtime_error("iknp correction length");
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
share[i] = (choices[i] & 1u) ? xor_block(masks[i], d[i]) : masks[i];
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/// @brief Sample this party's dealer pads. `me` is 0 or 1. Both parties pass
|
||||
/// the same lengths and call this before any other traffic on `link`.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline material sample(net::channel & link, int me, std::size_t nblock,
|
||||
std::size_t nbit, std::size_t nb2a, std::size_t ncw)
|
||||
{
|
||||
if (me != 0 && me != 1)
|
||||
throw std::invalid_argument("iknp party");
|
||||
|
||||
auto rand_bit = [] {
|
||||
return static_cast<std::uint8_t>(dpf::uniform_sample<unsigned char>() & 1u);
|
||||
};
|
||||
|
||||
std::vector<std::uint8_t> block_a(nblock), bit_a(nbit), bit_b(nbit), cw_bit(ncw), b2a_r(nb2a);
|
||||
std::vector<simde__m128i> block_b(nblock), cw_rand(ncw);
|
||||
for (std::size_t i = 0; i < nblock; ++i)
|
||||
{
|
||||
block_a[i] = rand_bit();
|
||||
block_b[i] = dpf::uniform_sample<simde__m128i>();
|
||||
}
|
||||
for (std::size_t i = 0; i < nbit; ++i)
|
||||
{
|
||||
bit_a[i] = rand_bit();
|
||||
bit_b[i] = rand_bit();
|
||||
}
|
||||
for (std::size_t i = 0; i < ncw; ++i)
|
||||
{
|
||||
cw_bit[i] = rand_bit();
|
||||
cw_rand[i] = dpf::uniform_sample<simde__m128i>();
|
||||
}
|
||||
for (std::size_t i = 0; i < nb2a; ++i)
|
||||
b2a_r[i] = rand_bit();
|
||||
|
||||
// Cross terms via two IKNP directions. Sender holds payload Δ, receiver
|
||||
// chooses χ; shares sum to χ·Δ.
|
||||
// D01 (P0 sends): χ=a1 / bit_a1 / cw_bit1, Δ=B0 / bit_b0 / cw_rand0.
|
||||
// D10 (P1 sends): χ=a0 / bit_a0 / cw_bit0, Δ=B1 / bit_b1 / cw_rand1.
|
||||
const std::size_t nxor = nblock + nbit + ncw;
|
||||
|
||||
std::vector<std::uint8_t> choices(nxor);
|
||||
std::vector<simde__m128i> payloads(nxor);
|
||||
for (std::size_t i = 0; i < nblock; ++i)
|
||||
{
|
||||
choices[i] = block_a[i];
|
||||
payloads[i] = block_b[i];
|
||||
}
|
||||
for (std::size_t j = 0; j < nbit; ++j)
|
||||
{
|
||||
choices[nblock + j] = bit_a[j];
|
||||
payloads[nblock + j] = detail::bit_block(bit_b[j]);
|
||||
}
|
||||
for (std::size_t k = 0; k < ncw; ++k)
|
||||
{
|
||||
choices[nblock + nbit + k] = cw_bit[k];
|
||||
payloads[nblock + nbit + k] = cw_rand[k];
|
||||
}
|
||||
|
||||
detail::role_state send_state; // used when we are IKNP sender (produce m0/m1)
|
||||
detail::role_state recv_state; // used when we are IKNP receiver
|
||||
std::vector<simde__m128i> send_m0, send_m1, recv_masks;
|
||||
|
||||
// D01: party 0 sends, party 1 receives.
|
||||
if (me == 0)
|
||||
detail::extend_send(link, send_state, nxor, send_m0, send_m1);
|
||||
else
|
||||
detail::extend_recv(link, recv_state, choices.data(), nxor, recv_masks);
|
||||
|
||||
std::vector<simde__m128i> d01_share;
|
||||
detail::correct_ot(link, me, me == 0,
|
||||
me == 0 ? send_m0 : std::vector<simde__m128i>{},
|
||||
me == 0 ? send_m1 : std::vector<simde__m128i>{},
|
||||
me == 0 ? payloads : std::vector<simde__m128i>{},
|
||||
me == 0 ? std::vector<simde__m128i>{} : recv_masks,
|
||||
me == 0 ? std::vector<std::uint8_t>{} : choices,
|
||||
d01_share);
|
||||
// d01_share_0 ⊕ d01_share_1 = χ1 · Δ0 = a1·B0 (etc.)
|
||||
|
||||
// D10: party 1 sends, party 0 receives.
|
||||
if (me == 0)
|
||||
detail::extend_recv(link, recv_state, choices.data(), nxor, recv_masks);
|
||||
else
|
||||
detail::extend_send(link, send_state, nxor, send_m0, send_m1);
|
||||
|
||||
std::vector<simde__m128i> d10_share;
|
||||
detail::correct_ot(link, me, me == 1,
|
||||
me == 1 ? send_m0 : std::vector<simde__m128i>{},
|
||||
me == 1 ? send_m1 : std::vector<simde__m128i>{},
|
||||
me == 1 ? payloads : std::vector<simde__m128i>{},
|
||||
me == 1 ? std::vector<simde__m128i>{} : recv_masks,
|
||||
me == 1 ? std::vector<std::uint8_t>{} : choices,
|
||||
d10_share);
|
||||
// d10_share_0 ⊕ d10_share_1 = χ0 · Δ1 = a0·B1 (etc.)
|
||||
|
||||
// Free large OT pads.
|
||||
send_m0.clear();
|
||||
send_m0.shrink_to_fit();
|
||||
send_m1.clear();
|
||||
send_m1.shrink_to_fit();
|
||||
recv_masks.clear();
|
||||
recv_masks.shrink_to_fit();
|
||||
|
||||
// CW gamma shares: gamma0 ⊕ gamma1 = (bit1·rand0) ⊕ (bit0·rand1)
|
||||
// = (D01 cw) ⊕ (D10 cw). Party 0 samples gamma0 and masks with both of
|
||||
// its OT shares; party 1 unmasks with both of its shares.
|
||||
std::vector<simde__m128i> gamma(ncw);
|
||||
if (ncw != 0)
|
||||
{
|
||||
const std::size_t cw_off = nblock + nbit;
|
||||
if (me == 0)
|
||||
{
|
||||
std::vector<simde__m128i> msg(ncw);
|
||||
for (std::size_t k = 0; k < ncw; ++k)
|
||||
{
|
||||
gamma[k] = dpf::uniform_sample<simde__m128i>();
|
||||
msg[k] = detail::xor_block(gamma[k],
|
||||
detail::xor_block(d01_share[cw_off + k], d10_share[cw_off + k]));
|
||||
}
|
||||
link.send_vec(msg, net::msg::bytes);
|
||||
}
|
||||
else
|
||||
{
|
||||
auto msg = link.recv_vec<simde__m128i>(net::msg::bytes);
|
||||
if (msg.size() != ncw)
|
||||
throw std::runtime_error("iknp cw pad length");
|
||||
for (std::size_t k = 0; k < ncw; ++k)
|
||||
gamma[k] = detail::xor_block(msg[k],
|
||||
detail::xor_block(d01_share[cw_off + k], d10_share[cw_off + k]));
|
||||
}
|
||||
}
|
||||
|
||||
// B2A: arithmetic shares of the XOR of the two random bits.
|
||||
std::vector<std::uint64_t> rho0(nb2a), arith_w(nb2a);
|
||||
if (nb2a != 0)
|
||||
{
|
||||
std::vector<simde__m128i> bm0, bm1, bmasks;
|
||||
if (me == 0)
|
||||
{
|
||||
detail::extend_send(link, send_state, nb2a, bm0, bm1);
|
||||
std::vector<std::uint64_t> corr(nb2a);
|
||||
for (std::size_t i = 0; i < nb2a; ++i)
|
||||
{
|
||||
rho0[i] = detail::low64(bm0[i]);
|
||||
const auto rho1 = detail::low64(bm1[i]);
|
||||
corr[i] = rho0[i] - rho1 + b2a_r[i];
|
||||
}
|
||||
link.send_vec(corr, net::msg::bytes);
|
||||
}
|
||||
else
|
||||
{
|
||||
detail::extend_recv(link, recv_state, b2a_r.data(), nb2a, bmasks);
|
||||
auto corr = link.recv_vec<std::uint64_t>(net::msg::bytes);
|
||||
if (corr.size() != nb2a)
|
||||
throw std::runtime_error("iknp b2a length");
|
||||
for (std::size_t i = 0; i < nb2a; ++i)
|
||||
{
|
||||
const auto rho = detail::low64(bmasks[i]);
|
||||
arith_w[i] = rho + static_cast<std::uint64_t>(b2a_r[i]) * corr[i];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
material out;
|
||||
out.blocks.resize(nblock);
|
||||
out.bits.resize(nbit);
|
||||
out.b2a.resize(nb2a);
|
||||
out.cws.resize(ncw);
|
||||
for (std::size_t i = 0; i < nblock; ++i)
|
||||
{
|
||||
// c0 ⊕ c1 = a0·B0 ⊕ a1·B1 ⊕ a1·B0 ⊕ a0·B1 = (a0⊕a1)·(B0⊕B1)
|
||||
const auto local = detail::gate(block_a[i], block_b[i]);
|
||||
// Party 0's share of a1·B0 is d01; of a0·B1 is d10. Same XOR for both
|
||||
// parties (their XOR shares already sum to the cross terms).
|
||||
const auto c = detail::xor_block(local,
|
||||
detail::xor_block(d01_share[i], d10_share[i]));
|
||||
out.blocks[i] = block_share{block_a[i], block_b[i], c};
|
||||
}
|
||||
for (std::size_t j = 0; j < nbit; ++j)
|
||||
{
|
||||
const auto off = nblock + j;
|
||||
const auto local = static_cast<std::uint8_t>(bit_a[j] & bit_b[j]);
|
||||
const auto c = static_cast<std::uint8_t>(local
|
||||
^ detail::lsb(d01_share[off]) ^ detail::lsb(d10_share[off]));
|
||||
out.bits[j] = bit_share{bit_a[j], bit_b[j], c};
|
||||
}
|
||||
for (std::size_t k = 0; k < ncw; ++k)
|
||||
out.cws[k] = cw_share{cw_rand[k], gamma[k], cw_bit[k]};
|
||||
for (std::size_t i = 0; i < nb2a; ++i)
|
||||
{
|
||||
std::uint64_t add;
|
||||
if (me == 0)
|
||||
add = static_cast<std::uint64_t>(b2a_r[i]) + (rho0[i] << 1);
|
||||
else
|
||||
add = static_cast<std::uint64_t>(b2a_r[i]) - (arith_w[i] << 1);
|
||||
out.b2a[i] = b2a_share{b2a_r[i], add};
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief 1-out-of-2 OT of 128-bit strings. The sender holds `(m0, m1)`. The
|
||||
/// receiver holds a choice bit per row and receives `m_choice`.
|
||||
/// @details Chou–Orlandi base OT runs on the first call for this `st`. Later
|
||||
/// calls only extend. `st` is one direction: the sender's state is
|
||||
/// not the receiver's state. `n == 0` sends nothing. Both parties
|
||||
/// pass the same `n`.
|
||||
/// @param ch peer channel
|
||||
/// @param me 0 or 1
|
||||
/// @param i_am_sender this party holds `m0` and `m1`
|
||||
/// @param st extension state for this direction
|
||||
/// @param m0 sender's first message, length `n` (ignored by the receiver)
|
||||
/// @param m1 sender's second message, length `n` (ignored by the receiver)
|
||||
/// @param choices receiver's choice bits, length `n` (ignored by the sender)
|
||||
/// @param got receiver's output `m_choice` (cleared for the sender)
|
||||
inline void transfer_labels(net::channel & ch, int me, bool i_am_sender,
|
||||
detail::role_state & st, const std::vector<simde__m128i> & m0,
|
||||
const std::vector<simde__m128i> & m1,
|
||||
const std::vector<std::uint8_t> & choices, std::vector<simde__m128i> & got)
|
||||
{
|
||||
if (me != 0 && me != 1)
|
||||
throw std::invalid_argument("iknp party");
|
||||
const std::size_t n = i_am_sender ? m0.size() : choices.size();
|
||||
if (i_am_sender)
|
||||
{
|
||||
if (m1.size() != n)
|
||||
throw std::invalid_argument("iknp transfer length");
|
||||
}
|
||||
else if (choices.size() != n)
|
||||
throw std::invalid_argument("iknp transfer length");
|
||||
got.clear();
|
||||
if (n == 0)
|
||||
return;
|
||||
if (i_am_sender)
|
||||
{
|
||||
std::vector<simde__m128i> r0, r1;
|
||||
detail::extend_send(ch, st, n, r0, r1);
|
||||
std::vector<simde__m128i> corr(2u * n);
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
corr[2u * i] = detail::xor_block(m0[i], r0[i]);
|
||||
corr[2u * i + 1u] = detail::xor_block(m1[i], r1[i]);
|
||||
}
|
||||
ch.send_vec(corr, net::msg::bytes);
|
||||
return;
|
||||
}
|
||||
std::vector<simde__m128i> masks;
|
||||
detail::extend_recv(ch, st, choices.data(), n, masks);
|
||||
const auto corr = ch.recv_vec<simde__m128i>(net::msg::bytes);
|
||||
if (corr.size() != 2u * n)
|
||||
throw std::runtime_error("iknp transfer correction length");
|
||||
got.resize(n);
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
const std::size_t slot = 2u * i + static_cast<std::size_t>(choices[i] & 1u);
|
||||
got[i] = detail::xor_block(masks[i], corr[slot]);
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Channel rounds inside `sample` (base OT, extension chunks, corrections).
|
||||
/// @details Two directions when `nblock+nbit+ncw > 0`: each is a 2-message
|
||||
/// Chou–Orlandi base, one U-matrix per `chunk_rows`, and one
|
||||
/// correction. CW gamma is one more round. B2A reuses the base OT
|
||||
/// and adds an extension plus a correction. Add this to
|
||||
/// `plan::rounds()` via `rounds_including`.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
HEDLEY_PURE
|
||||
inline std::size_t setup_rounds(std::size_t nblock, std::size_t nbit,
|
||||
std::size_t nb2a, std::size_t ncw) noexcept
|
||||
{
|
||||
const auto chunks = [](std::size_t n) {
|
||||
return n == 0 ? std::size_t{0}
|
||||
: (n + detail::chunk_rows - 1) / detail::chunk_rows;
|
||||
};
|
||||
const std::size_t nxor = nblock + nbit + ncw;
|
||||
std::size_t rounds = 0;
|
||||
if (nxor != 0)
|
||||
{
|
||||
rounds += 2 + chunks(nxor) + 1;
|
||||
rounds += 2 + chunks(nxor) + 1;
|
||||
}
|
||||
if (ncw != 0)
|
||||
++rounds;
|
||||
if (nb2a != 0)
|
||||
rounds += chunks(nb2a) + 1;
|
||||
return rounds;
|
||||
}
|
||||
|
||||
} // namespace iknp
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_IKNP_HPP__
|
||||
32
include/dpf/iknp_graphs.hpp
Normal file
32
include/dpf/iknp_graphs.hpp
Normal file
|
|
@ -0,0 +1,32 @@
|
|||
/// @file dpf/iknp_graphs.hpp
|
||||
/// @brief IKNP setup as `schedule_round` lists (pulls Asio; keep out of app_plans).
|
||||
#ifndef LIBDPF_INCLUDE_DPF_IKNP_GRAPHS_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_IKNP_GRAPHS_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <memory>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/iknp.hpp"
|
||||
#include "dpf/pad_graphs.hpp"
|
||||
#include "dpf/protocol.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace protocol
|
||||
{
|
||||
|
||||
/// @brief IKNP `sample` wire frames as schedule rounds (count from setup_rounds).
|
||||
inline std::vector<schedule_round> iknp_setup_graph(std::size_t nblock,
|
||||
std::size_t nbit, std::size_t nb2a, std::size_t ncw,
|
||||
const std::shared_ptr<std::vector<std::uint8_t>> & tape)
|
||||
{
|
||||
return make_pad_rounds(iknp::setup_rounds(nblock, nbit, nb2a, ncw), 16, tape,
|
||||
edge_channel::peer);
|
||||
}
|
||||
|
||||
} // namespace protocol
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_IKNP_GRAPHS_HPP__
|
||||
File diff suppressed because it is too large
Load diff
568
include/dpf/interleave_leaves.hpp
Normal file
568
include/dpf/interleave_leaves.hpp
Normal file
|
|
@ -0,0 +1,568 @@
|
|||
/// @file dpf/interleave_leaves.hpp
|
||||
/// @brief Interleave (and deinterleave) per-key leaf vectors into cohort order.
|
||||
/// @details Given `nkeys` vectors of length `nleaves`, write one buffer whose
|
||||
/// logical lane index is `cohort_index(i, k, nkeys) = i * nkeys + k`
|
||||
/// (leaf `i` of key `k`). That matches sequence-cohort layout and the
|
||||
/// key-major axis of interval-cohort leaves (lanes inside a leaf stay
|
||||
/// contiguous in each one-key buffer; this routine interleaves whole
|
||||
/// leaf *values*, not interior walk nodes).
|
||||
///
|
||||
/// Storage and packing:
|
||||
/// - `uint64_t` / `uint32_t` / `uint16_t` / `uint8_t`: one array element
|
||||
/// per leaf; `out[i * nkeys + k] = keys[k][i]`.
|
||||
/// - `dpf::bit`, `dpf::twobit`, `dpf::nyble`, and the matching
|
||||
/// `dpf::gf2` / `gf22` / `gf24` lanes: same logical order, packed
|
||||
/// tightly into `uint64_t` words like `dynamic_bit_array` /
|
||||
/// `dynamic_packed_array` — lane `j` occupies bits
|
||||
/// `[j * W, (j + 1) * W)` of the flat stream (`W` = 1, 2, or 4),
|
||||
/// low lane in the low bits of each word. So a naive *element*
|
||||
/// index `i * nkeys + k` is wrong for sub-byte widths; use the
|
||||
/// logical index above, then pack with width `W`.
|
||||
///
|
||||
/// Kernels are word-wise (and 8×8 bit transpose batches for 1-bit),
|
||||
/// not a scalar loop per bit. Deinterleave is the inverse transpose.
|
||||
/// @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_INTERLEAVE_LEAVES_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_INTERLEAVE_LEAVES_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <type_traits>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/bit.hpp"
|
||||
#include "dpf/modint.hpp"
|
||||
#include "dpf/nyble.hpp"
|
||||
#include "dpf/twobit.hpp"
|
||||
#include "dpf/utils.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
/// @brief Underlying word/element type used to store a sequence of `LaneT` leaves.
|
||||
template <typename LaneT, typename Enable = void>
|
||||
struct leaf_storage
|
||||
{
|
||||
using type = LaneT;
|
||||
};
|
||||
|
||||
template <typename LaneT>
|
||||
struct leaf_storage<LaneT, std::enable_if_t<utils::is_packed_subbyte_v<LaneT>>>
|
||||
{
|
||||
using type = std::uint64_t;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct leaf_storage<dpf::bit>
|
||||
{
|
||||
using type = std::uint64_t;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct leaf_storage<dpf::twobit>
|
||||
{
|
||||
using type = std::uint64_t;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct leaf_storage<dpf::nyble>
|
||||
{
|
||||
using type = std::uint64_t;
|
||||
};
|
||||
|
||||
template <typename LaneT>
|
||||
using leaf_storage_t = typename leaf_storage<LaneT>::type;
|
||||
|
||||
/// @brief Number of `leaf_storage_t<LaneT>` words needed for `n` logical leaves.
|
||||
template <typename LaneT>
|
||||
HEDLEY_CONST
|
||||
HEDLEY_NO_THROW
|
||||
constexpr std::size_t leaf_storage_words(std::size_t n) noexcept
|
||||
{
|
||||
if constexpr (utils::is_packed_subbyte_v<LaneT>)
|
||||
{
|
||||
constexpr std::size_t w = utils::packed_lane_bits_v<LaneT>;
|
||||
return (n * w + 63u) / 64u;
|
||||
}
|
||||
else
|
||||
{
|
||||
return n;
|
||||
}
|
||||
}
|
||||
|
||||
namespace interleave_detail
|
||||
{
|
||||
|
||||
/// @brief 8×8 bit matrix transpose in a 64-bit word (Hacker's Delight).
|
||||
/// @details Byte `r` holds row `r`; bit `c` of that byte is column `c`
|
||||
/// (LSB = column 0). Transpose swaps rows and columns.
|
||||
HEDLEY_CONST
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
constexpr std::uint64_t bit_transpose_8x8(std::uint64_t x) noexcept
|
||||
{
|
||||
std::uint64_t t = (x ^ (x >> 7)) & 0x00AA00AA00AA00AAull;
|
||||
x = x ^ t ^ (t << 7);
|
||||
t = (x ^ (x >> 14)) & 0x0000CCCC0000CCCCull;
|
||||
x = x ^ t ^ (t << 14);
|
||||
t = (x ^ (x >> 28)) & 0x00000000F0F0F0F0ull;
|
||||
x = x ^ t ^ (t << 28);
|
||||
return x;
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void write_bits(std::uint64_t * out, std::size_t bit_pos, std::uint64_t bits,
|
||||
unsigned nbits) noexcept
|
||||
{
|
||||
if (nbits == 0)
|
||||
return;
|
||||
const std::size_t word = bit_pos / 64u;
|
||||
const unsigned shift = static_cast<unsigned>(bit_pos % 64u);
|
||||
const std::uint64_t mask = nbits == 64
|
||||
? ~std::uint64_t{0}
|
||||
: ((std::uint64_t{1} << nbits) - 1u);
|
||||
const std::uint64_t val = bits & mask;
|
||||
out[word] |= val << shift;
|
||||
if (shift + nbits > 64u)
|
||||
out[word + 1u] |= val >> (64u - shift);
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
std::uint64_t read_bits(const std::uint64_t * in, std::size_t bit_pos,
|
||||
unsigned nbits) noexcept
|
||||
{
|
||||
if (nbits == 0)
|
||||
return 0;
|
||||
const std::size_t word = bit_pos / 64u;
|
||||
const unsigned shift = static_cast<unsigned>(bit_pos % 64u);
|
||||
const std::uint64_t mask = nbits == 64
|
||||
? ~std::uint64_t{0}
|
||||
: ((std::uint64_t{1} << nbits) - 1u);
|
||||
std::uint64_t val = in[word] >> shift;
|
||||
if (shift + nbits > 64u)
|
||||
val |= in[word + 1u] << (64u - shift);
|
||||
return val & mask;
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
std::uint8_t load_key_byte_bits(const std::uint64_t * key, std::size_t nleaves,
|
||||
std::size_t i0) noexcept
|
||||
{
|
||||
if (i0 >= nleaves)
|
||||
return 0;
|
||||
const unsigned take = static_cast<unsigned>(
|
||||
nleaves - i0 < 8u ? nleaves - i0 : 8u);
|
||||
return static_cast<std::uint8_t>(read_bits(key, i0, take));
|
||||
}
|
||||
|
||||
/// @brief Interleave 1-bit lanes with 8×8 transpose batches plus a bit writer.
|
||||
HEDLEY_NO_THROW
|
||||
inline void interleave_bits(std::uint64_t * HEDLEY_RESTRICT out,
|
||||
const std::uint64_t * const * HEDLEY_RESTRICT keys, std::size_t nkeys,
|
||||
std::size_t nleaves) noexcept
|
||||
{
|
||||
const std::size_t out_words = leaf_storage_words<dpf::bit>(nkeys * nleaves);
|
||||
if (out_words != 0)
|
||||
std::memset(out, 0, out_words * sizeof(std::uint64_t));
|
||||
if (nkeys == 0 || nleaves == 0)
|
||||
return;
|
||||
|
||||
for (std::size_t i0 = 0; i0 < nleaves; i0 += 8u)
|
||||
{
|
||||
const unsigned nrows = static_cast<unsigned>(
|
||||
nleaves - i0 < 8u ? nleaves - i0 : 8u);
|
||||
for (std::size_t k0 = 0; k0 < nkeys; k0 += 8u)
|
||||
{
|
||||
const unsigned ncols = static_cast<unsigned>(
|
||||
nkeys - k0 < 8u ? nkeys - k0 : 8u);
|
||||
std::uint64_t packed = 0;
|
||||
DPF_UNROLL_LOOP
|
||||
for (unsigned c = 0; c < 8u; ++c)
|
||||
{
|
||||
std::uint8_t b = 0;
|
||||
if (c < ncols)
|
||||
b = load_key_byte_bits(keys[k0 + c], nleaves, i0);
|
||||
packed |= static_cast<std::uint64_t>(b) << (8u * c);
|
||||
}
|
||||
const std::uint64_t t = bit_transpose_8x8(packed);
|
||||
DPF_UNROLL_LOOP
|
||||
for (unsigned r = 0; r < nrows; ++r)
|
||||
{
|
||||
const std::uint64_t row
|
||||
= (t >> (8u * r)) & ((std::uint64_t{1} << ncols) - 1u);
|
||||
write_bits(out, (i0 + r) * nkeys + k0, row, ncols);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
inline void deinterleave_bits(std::uint64_t * const * HEDLEY_RESTRICT keys,
|
||||
const std::uint64_t * HEDLEY_RESTRICT in, std::size_t nkeys,
|
||||
std::size_t nleaves) noexcept
|
||||
{
|
||||
if (nkeys == 0 || nleaves == 0)
|
||||
return;
|
||||
for (std::size_t k = 0; k < nkeys; ++k)
|
||||
{
|
||||
const std::size_t words = leaf_storage_words<dpf::bit>(nleaves);
|
||||
if (words != 0)
|
||||
std::memset(keys[k], 0, words * sizeof(std::uint64_t));
|
||||
}
|
||||
|
||||
for (std::size_t i0 = 0; i0 < nleaves; i0 += 8u)
|
||||
{
|
||||
const unsigned nrows = static_cast<unsigned>(
|
||||
nleaves - i0 < 8u ? nleaves - i0 : 8u);
|
||||
for (std::size_t k0 = 0; k0 < nkeys; k0 += 8u)
|
||||
{
|
||||
const unsigned ncols = static_cast<unsigned>(
|
||||
nkeys - k0 < 8u ? nkeys - k0 : 8u);
|
||||
std::uint64_t packed = 0;
|
||||
DPF_UNROLL_LOOP
|
||||
for (unsigned r = 0; r < nrows; ++r)
|
||||
{
|
||||
const std::uint64_t row
|
||||
= read_bits(in, (i0 + r) * nkeys + k0, ncols);
|
||||
packed |= row << (8u * r);
|
||||
}
|
||||
// Inverse of the interleave transpose: same 8×8 transpose.
|
||||
const std::uint64_t t = bit_transpose_8x8(packed);
|
||||
DPF_UNROLL_LOOP
|
||||
for (unsigned c = 0; c < ncols; ++c)
|
||||
{
|
||||
const std::uint64_t col
|
||||
= (t >> (8u * c)) & ((std::uint64_t{1} << nrows) - 1u);
|
||||
write_bits(keys[k0 + c], i0, col, nrows);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Interleave `W`-bit lanes (`W` = 2 or 4) by filling output words.
|
||||
template <unsigned W>
|
||||
HEDLEY_NO_THROW
|
||||
void interleave_packed_w(std::uint64_t * HEDLEY_RESTRICT out,
|
||||
const std::uint64_t * const * HEDLEY_RESTRICT keys, std::size_t nkeys,
|
||||
std::size_t nleaves) noexcept
|
||||
{
|
||||
static_assert(W == 2 || W == 4, "packed width must be 2 or 4");
|
||||
constexpr std::uint64_t lane_mask = (std::uint64_t{1} << W) - 1u;
|
||||
constexpr unsigned lanes_per_word = 64u / W;
|
||||
const std::size_t total = nkeys * nleaves;
|
||||
const std::size_t out_words = (total * W + 63u) / 64u;
|
||||
if (out_words != 0)
|
||||
std::memset(out, 0, out_words * sizeof(std::uint64_t));
|
||||
if (nkeys == 0 || nleaves == 0)
|
||||
return;
|
||||
|
||||
for (std::size_t ow = 0; ow < out_words; ++ow)
|
||||
{
|
||||
std::uint64_t word = 0;
|
||||
const std::size_t base = ow * lanes_per_word;
|
||||
const unsigned nlanes = static_cast<unsigned>(
|
||||
total - base < lanes_per_word ? total - base : lanes_per_word);
|
||||
DPF_UNROLL_LOOP
|
||||
for (unsigned t = 0; t < nlanes; ++t)
|
||||
{
|
||||
const std::size_t g = base + t;
|
||||
const std::size_t i = g / nkeys;
|
||||
const std::size_t k = g - i * nkeys;
|
||||
const std::size_t src_bit = i * W;
|
||||
const std::uint64_t lane
|
||||
= (keys[k][src_bit / 64u] >> (src_bit % 64u)) & lane_mask;
|
||||
word |= lane << (t * W);
|
||||
}
|
||||
out[ow] = word;
|
||||
}
|
||||
}
|
||||
|
||||
template <unsigned W>
|
||||
HEDLEY_NO_THROW
|
||||
void deinterleave_packed_w(std::uint64_t * const * HEDLEY_RESTRICT keys,
|
||||
const std::uint64_t * HEDLEY_RESTRICT in, std::size_t nkeys,
|
||||
std::size_t nleaves) noexcept
|
||||
{
|
||||
static_assert(W == 2 || W == 4, "packed width must be 2 or 4");
|
||||
constexpr std::uint64_t lane_mask = (std::uint64_t{1} << W) - 1u;
|
||||
if (nkeys == 0 || nleaves == 0)
|
||||
return;
|
||||
for (std::size_t k = 0; k < nkeys; ++k)
|
||||
{
|
||||
const std::size_t words = (nleaves * W + 63u) / 64u;
|
||||
if (words != 0)
|
||||
std::memset(keys[k], 0, words * sizeof(std::uint64_t));
|
||||
}
|
||||
|
||||
const std::size_t total = nkeys * nleaves;
|
||||
for (std::size_t g = 0; g < total; ++g)
|
||||
{
|
||||
const std::size_t i = g / nkeys;
|
||||
const std::size_t k = g - i * nkeys;
|
||||
const std::size_t src_bit = g * W;
|
||||
const std::uint64_t lane
|
||||
= (in[src_bit / 64u] >> (src_bit % 64u)) & lane_mask;
|
||||
const std::size_t dst_bit = i * W;
|
||||
keys[k][dst_bit / 64u] |= lane << (dst_bit % 64u);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
HEDLEY_NO_THROW
|
||||
void interleave_wide(T * HEDLEY_RESTRICT out, const T * const * HEDLEY_RESTRICT keys,
|
||||
std::size_t nkeys, std::size_t nleaves) noexcept
|
||||
{
|
||||
if (nkeys == 0 || nleaves == 0)
|
||||
return;
|
||||
if (nkeys == 1)
|
||||
{
|
||||
std::memcpy(out, keys[0], nleaves * sizeof(T));
|
||||
return;
|
||||
}
|
||||
for (std::size_t i = 0; i < nleaves; ++i)
|
||||
{
|
||||
T * dst = out + i * nkeys;
|
||||
DPF_UNROLL_LOOP
|
||||
for (std::size_t k = 0; k < nkeys; ++k)
|
||||
dst[k] = keys[k][i];
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
HEDLEY_NO_THROW
|
||||
void deinterleave_wide(T * const * HEDLEY_RESTRICT keys,
|
||||
const T * HEDLEY_RESTRICT in, std::size_t nkeys, std::size_t nleaves) noexcept
|
||||
{
|
||||
if (nkeys == 0 || nleaves == 0)
|
||||
return;
|
||||
if (nkeys == 1)
|
||||
{
|
||||
std::memcpy(keys[0], in, nleaves * sizeof(T));
|
||||
return;
|
||||
}
|
||||
for (std::size_t i = 0; i < nleaves; ++i)
|
||||
{
|
||||
const T * src = in + i * nkeys;
|
||||
DPF_UNROLL_LOOP
|
||||
for (std::size_t k = 0; k < nkeys; ++k)
|
||||
keys[k][i] = src[k];
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace interleave_detail
|
||||
|
||||
/// @brief Interleave `nkeys` leaf vectors of length `nleaves` into `out`.
|
||||
/// @tparam LaneT an integer, a packed lane (`bit`, `twobit`, `nyble`,
|
||||
/// `gf2`, `gf22`, `gf24`), or a trivially copyable element of at
|
||||
/// most 8 bytes (`gf28`, `gf216`, `gf232`, `gf264`)
|
||||
/// @param out destination storage (`leaf_storage_t<LaneT>`; see file comment)
|
||||
/// @param keys array of `nkeys` pointers to per-key leaf storage
|
||||
/// @param nkeys number of keys (`m`)
|
||||
/// @param nleaves number of leaves per key (`L`)
|
||||
/// \complexity Θ(`nkeys * nleaves`) lane moves. Wide types stream by leaf;
|
||||
/// 1-bit uses 8×8 transpose tiles; 2/4-bit fill output words.
|
||||
template <typename LaneT>
|
||||
HEDLEY_NO_THROW
|
||||
void interleave_leaves(leaf_storage_t<LaneT> * out,
|
||||
const leaf_storage_t<LaneT> * const * keys, std::size_t nkeys,
|
||||
std::size_t nleaves) noexcept
|
||||
{
|
||||
if constexpr (utils::is_packed_subbyte_v<LaneT>)
|
||||
{
|
||||
constexpr unsigned w = utils::packed_lane_bits_v<LaneT>;
|
||||
if constexpr (w == 1)
|
||||
interleave_detail::interleave_bits(out, keys, nkeys, nleaves);
|
||||
else if constexpr (w == 2)
|
||||
interleave_detail::interleave_packed_w<2>(out, keys, nkeys, nleaves);
|
||||
else
|
||||
{
|
||||
static_assert(w == 4, "packed lane width must be 1, 2, or 4");
|
||||
interleave_detail::interleave_packed_w<4>(out, keys, nkeys, nleaves);
|
||||
}
|
||||
}
|
||||
else if constexpr ((std::is_integral_v<LaneT> && !std::is_same_v<LaneT, bool>)
|
||||
|| (std::is_trivially_copyable_v<LaneT> && std::is_standard_layout_v<LaneT>
|
||||
&& sizeof(LaneT) > 0 && sizeof(LaneT) <= 8))
|
||||
{
|
||||
interleave_detail::interleave_wide(out, keys, nkeys, nleaves);
|
||||
}
|
||||
else
|
||||
{
|
||||
static_assert(std::is_integral_v<LaneT>,
|
||||
"interleave_leaves LaneT must be an integer, a packed lane, or a field element of at most 8 bytes");
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Inverse of `interleave_leaves`: split one interleaved buffer into
|
||||
/// `nkeys` per-key leaf vectors of length `nleaves`.
|
||||
/// @tparam LaneT same set as `interleave_leaves`
|
||||
/// @param keys destination per-key storage pointers
|
||||
/// @param in interleaved source
|
||||
/// @param nkeys number of keys
|
||||
/// @param nleaves number of leaves per key
|
||||
/// \complexity Same order as `interleave_leaves`.
|
||||
template <typename LaneT>
|
||||
HEDLEY_NO_THROW
|
||||
void deinterleave_leaves(leaf_storage_t<LaneT> * const * keys,
|
||||
const leaf_storage_t<LaneT> * in, std::size_t nkeys,
|
||||
std::size_t nleaves) noexcept
|
||||
{
|
||||
if constexpr (utils::is_packed_subbyte_v<LaneT>)
|
||||
{
|
||||
constexpr unsigned w = utils::packed_lane_bits_v<LaneT>;
|
||||
if constexpr (w == 1)
|
||||
interleave_detail::deinterleave_bits(keys, in, nkeys, nleaves);
|
||||
else if constexpr (w == 2)
|
||||
interleave_detail::deinterleave_packed_w<2>(keys, in, nkeys, nleaves);
|
||||
else
|
||||
{
|
||||
static_assert(w == 4, "packed lane width must be 1, 2, or 4");
|
||||
interleave_detail::deinterleave_packed_w<4>(keys, in, nkeys, nleaves);
|
||||
}
|
||||
}
|
||||
else if constexpr ((std::is_integral_v<LaneT> && !std::is_same_v<LaneT, bool>)
|
||||
|| (std::is_trivially_copyable_v<LaneT> && std::is_standard_layout_v<LaneT>
|
||||
&& sizeof(LaneT) > 0 && sizeof(LaneT) <= 8))
|
||||
{
|
||||
interleave_detail::deinterleave_wide(keys, in, nkeys, nleaves);
|
||||
}
|
||||
else
|
||||
{
|
||||
static_assert(std::is_integral_v<LaneT>,
|
||||
"deinterleave_leaves LaneT must be an integer, a packed lane, or a field element of at most 8 bytes");
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Dot an interleaved 1-bit cohort with `weights`, as `modint<Nbits>`.
|
||||
/// @details `bits` is the buffer from `interleave_leaves<dpf::bit>`: leaf `i`
|
||||
/// is the integer whose bit `k` is key `k`'s leaf `i`. That integer,
|
||||
/// reduced modulo `2^Nbits` (the low `Nbits` bits), is one `modint`.
|
||||
/// The result is `sum_i modint(leaf i) * weights[i]`. `weights[i]`
|
||||
/// may be a `modint<Nbits>` or an integer.
|
||||
/// When `nkeys == Nbits` and that width is 8, 16, 32, 64, or 128,
|
||||
/// the leaves are contiguous native words and the product is a
|
||||
/// straight multiply-accumulate.
|
||||
/// @tparam Nbits modulus width, `1` through `256`
|
||||
/// @param bits interleaved 1-bit stream
|
||||
/// @param nleaves number of leaves (domain points)
|
||||
/// @param nkeys number of keys that were interleaved
|
||||
/// @param weights one weight per leaf
|
||||
/// @return the inner product in `modint<Nbits>`
|
||||
/// \complexity Θ(`nleaves`) multiplications in the `modint` word. Aligned
|
||||
/// widths do one native multiply per leaf; other widths extract
|
||||
/// the `Nbits` bits of each leaf first.
|
||||
template <std::size_t Nbits, typename Weights>
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
modint<Nbits> interleaved_bits_inner_product(const std::uint64_t * bits,
|
||||
std::size_t nleaves, std::size_t nkeys, Weights && weights)
|
||||
{
|
||||
using mod = modint<Nbits>;
|
||||
using limb = typename mod::integral_type;
|
||||
mod acc{};
|
||||
if (nleaves == 0 || nkeys == 0 || bits == nullptr)
|
||||
return acc;
|
||||
|
||||
auto limb_of = [](auto v) -> limb {
|
||||
using T = std::decay_t<decltype(v)>;
|
||||
if constexpr (std::is_same_v<T, mod>)
|
||||
return static_cast<limb>(v);
|
||||
else
|
||||
return static_cast<limb>(mod(static_cast<limb>(v)));
|
||||
};
|
||||
|
||||
auto mac = [&](limb v, std::size_t i) {
|
||||
acc += mod(v) * mod(limb_of(weights[i]));
|
||||
};
|
||||
|
||||
if constexpr (Nbits == 128)
|
||||
{
|
||||
if (nkeys == 128)
|
||||
{
|
||||
limb sum{};
|
||||
for (std::size_t i = 0; i < nleaves; ++i)
|
||||
{
|
||||
const std::uint64_t * p = bits + i * 2u;
|
||||
const limb v = static_cast<limb>(p[0])
|
||||
| (static_cast<limb>(p[1]) << 64);
|
||||
sum += v * limb_of(weights[i]);
|
||||
}
|
||||
return mod(sum);
|
||||
}
|
||||
}
|
||||
if constexpr (Nbits == 8 || Nbits == 16 || Nbits == 32 || Nbits == 64)
|
||||
{
|
||||
if (nkeys == Nbits)
|
||||
{
|
||||
limb sum{};
|
||||
for (std::size_t i = 0; i < nleaves; ++i)
|
||||
{
|
||||
limb v{};
|
||||
if constexpr (Nbits == 64)
|
||||
v = static_cast<limb>(bits[i]);
|
||||
else if constexpr (Nbits == 32)
|
||||
{
|
||||
const std::uint64_t word = bits[i / 2u];
|
||||
v = static_cast<limb>((i & 1u) ? (word >> 32) : (word & 0xffffffffu));
|
||||
}
|
||||
else if constexpr (Nbits == 16)
|
||||
{
|
||||
const std::uint64_t word = bits[i / 4u];
|
||||
v = static_cast<limb>((word >> ((i % 4u) * 16u)) & 0xffffu);
|
||||
}
|
||||
else
|
||||
{
|
||||
const std::uint64_t word = bits[i / 8u];
|
||||
v = static_cast<limb>((word >> ((i % 8u) * 8u)) & 0xffu);
|
||||
}
|
||||
sum += v * limb_of(weights[i]);
|
||||
}
|
||||
return mod(sum);
|
||||
}
|
||||
}
|
||||
|
||||
for (std::size_t i = 0; i < nleaves; ++i)
|
||||
{
|
||||
const std::size_t bit_pos = i * nkeys;
|
||||
limb v{};
|
||||
if constexpr (Nbits <= 64)
|
||||
{
|
||||
const unsigned take = static_cast<unsigned>(
|
||||
nkeys < Nbits ? nkeys : Nbits);
|
||||
v = static_cast<limb>(interleave_detail::read_bits(bits, bit_pos, take));
|
||||
}
|
||||
else
|
||||
{
|
||||
unsigned left = static_cast<unsigned>(Nbits);
|
||||
std::size_t pos = bit_pos;
|
||||
unsigned shift = 0;
|
||||
while (left != 0 && shift < nkeys)
|
||||
{
|
||||
const unsigned room = static_cast<unsigned>(nkeys - shift);
|
||||
const unsigned n = left < 64u ? left : 64u;
|
||||
const unsigned take = n < room ? n : room;
|
||||
const auto chunk = interleave_detail::read_bits(bits, pos, take);
|
||||
v |= static_cast<limb>(chunk) << shift;
|
||||
shift += take;
|
||||
pos += take;
|
||||
left -= take;
|
||||
if (take == 0)
|
||||
break;
|
||||
}
|
||||
}
|
||||
mac(v, i);
|
||||
}
|
||||
return acc;
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_INTERLEAVE_LEAVES_HPP__
|
||||
|
|
@ -5,7 +5,9 @@
|
|||
/// `p ≤ (x − r) mod 2^n ≤ q`, and the false payload otherwise.
|
||||
///
|
||||
/// The key is one `lt` comparison at `γ = r − 1`, the Boyle–Chandran–
|
||||
/// Gilboa–Gupta–Ishai–Kumar–Rathee reduction (EUROCRYPT 2021, Fig. 3).
|
||||
/// Gilboa–Gupta–Ishai–Kumar–Rathee reduction (EUROCRYPT 2021, Fig. 3;
|
||||
/// ePrint 2020/1392). Their Section 4.1 is one DCF for a public interval,
|
||||
/// where the earlier gate used about two.
|
||||
/// Evaluation walks that key at the two public shifts of `x` and adds
|
||||
/// a secret-shared correction. Seed corrections, advice bits, leaves,
|
||||
/// and the path memoizer stay single-path.
|
||||
|
|
@ -26,6 +28,8 @@
|
|||
#include "dpf/dcf.hpp"
|
||||
#include "dpf/eval_unified.hpp"
|
||||
#include "dpf/geneval.hpp"
|
||||
#include "dpf/grow.hpp"
|
||||
#include "dpf/grow_ds.hpp"
|
||||
#include "dpf/incremental.hpp"
|
||||
#include "dpf/output_buffer.hpp"
|
||||
#include "dpf/secret_share.hpp"
|
||||
|
|
@ -86,7 +90,9 @@ struct ic_key
|
|||
detail::cmp_group_info<dpf::concrete_type_t<Beta>>::custom,
|
||||
detail::group_elem, uint64_t>;
|
||||
|
||||
key_type key;
|
||||
/// @brief Inner comparison key. Eval that accepts a `dpf_key` also accepts
|
||||
/// this object and reads `dpf_key`.
|
||||
key_type dpf_key;
|
||||
uint64_t lo = 0;
|
||||
uint64_t hi = 0;
|
||||
uint64_t input_mask = 0;
|
||||
|
|
@ -103,7 +109,7 @@ struct ic_key
|
|||
ic_key(key_type k, uint64_t lo_in, uint64_t hi_in, uint64_t nmask,
|
||||
uint64_t gmask, share_type dshare, share_type cshare, share_type dcoeff,
|
||||
share_type ccoeff) noexcept(std::is_nothrow_move_constructible_v<key_type>)
|
||||
: key(std::move(k))
|
||||
: dpf_key(std::move(k))
|
||||
, lo(lo_in)
|
||||
, hi(hi_in)
|
||||
, input_mask(nmask)
|
||||
|
|
@ -417,7 +423,8 @@ Input gamma_of(Input r)
|
|||
}
|
||||
|
||||
template <typename IcKey, typename Query, typename Memo>
|
||||
auto eval_one(const IcKey & k, Query && x, Memo & memo)
|
||||
auto eval_one(const IcKey & k, Query && x, Memo & memo,
|
||||
proof_token * pi = nullptr)
|
||||
{
|
||||
if (!k.assigned)
|
||||
throw std::invalid_argument(
|
||||
|
|
@ -435,10 +442,10 @@ auto eval_one(const IcKey & k, Query && x, Memo & memo)
|
|||
else
|
||||
return detail::group_from_beta(v);
|
||||
};
|
||||
const auto a = opened(eval_point<beta>(dpf::cmp, k.key,
|
||||
input_from_bits<in_type>(xp), memo));
|
||||
const auto b = opened(eval_point<beta>(dpf::cmp, k.key,
|
||||
input_from_bits<in_type>(xq), memo));
|
||||
const auto a = opened(detail::incr::eval_cmp_point_impl<beta>(k.dpf_key,
|
||||
input_from_bits<in_type>(xp), memo, pi));
|
||||
const auto b = opened(detail::incr::eval_cmp_point_impl<beta>(k.dpf_key,
|
||||
input_from_bits<in_type>(xq), memo, pi));
|
||||
const int cx = public_cx(xu, k.lo, k.hi, k.input_mask);
|
||||
auto scaled = detail::group_zero(a);
|
||||
if (cx == 1)
|
||||
|
|
@ -453,10 +460,12 @@ auto eval_one(const IcKey & k, Query && x, Memo & memo)
|
|||
else
|
||||
{
|
||||
const uint64_t a = opened_u64(
|
||||
eval_point(dpf::cmp, k.key, input_from_bits<in_type>(xp), memo),
|
||||
detail::incr::eval_cmp_point_impl(k.dpf_key,
|
||||
input_from_bits<in_type>(xp), memo, pi),
|
||||
k.group_mask);
|
||||
const uint64_t b = opened_u64(
|
||||
eval_point(dpf::cmp, k.key, input_from_bits<in_type>(xq), memo),
|
||||
detail::incr::eval_cmp_point_impl(k.dpf_key,
|
||||
input_from_bits<in_type>(xq), memo, pi),
|
||||
k.group_mask);
|
||||
const int cx = public_cx(xu, k.lo, k.hi, k.input_mask);
|
||||
uint64_t scaled = 0;
|
||||
|
|
@ -482,6 +491,8 @@ auto eval_one(const IcKey & k, Query && x, Memo & memo)
|
|||
/// @param r the secret input mask
|
||||
/// @param spec the public bounds and payloads
|
||||
/// @return Dealer key for public bounds `spec` and secret mask `r`
|
||||
/// @note Following Boyle, Chandran, Gilboa, Gupta, Ishai, Kumar, and Rathee, EUROCRYPT 2021, Fig. 3 (ePrint 2020/1392): one comparison key, evaluated at two public shifts.
|
||||
/// \complexity O(n) time and O(n) key size. n is the domain bitlength (`depth`). The loop does two interior PRG expansions and writes one correction word per level, then builds one exterior leaf per output.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
|
|
@ -516,16 +527,23 @@ auto make_dpf(InputT && r, const ic_pack<Beta> & spec)
|
|||
/// @param r1 party 1's share of the mask
|
||||
/// @param rng the Doerner–Shelat randomness tapes
|
||||
/// @param spec the public bounds and payloads
|
||||
/// @param tags optional `verifiable` or `extractable` markers
|
||||
/// @return the two party keys
|
||||
/// \complexity O(n) time. One `ds_advance_level` per level: two PRG expansions and one `prepare_level`. n is `depth`.
|
||||
/// \rounds No sockets. This is the in-process transcript. A networked walk is `dpf::party::dist::point_party`.
|
||||
/// \communication none here. `local_cw_protocol` opens the correction word locally.
|
||||
/// \preprocessing Per level, `prepare_level` draws one `ds_cw_pads` (two parties × a 128-bit rand, a 128-bit gamma, and a bit) and two `ds_and_pads`. Arithmetic inputs also run a carry chain of n-1 bit-AND triples in `encode_walk_shares`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename RootSampler,
|
||||
typename PadRng,
|
||||
typename InputT,
|
||||
typename Beta>
|
||||
typename Beta,
|
||||
typename ...Tags>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto make_dpf_doerner_shelat(InputT r0, InputT r1,
|
||||
ds_randomness<RootSampler, PadRng> rng, const ic_pack<Beta> & spec)
|
||||
ds_randomness<RootSampler, PadRng> rng, const ic_pack<Beta> & spec,
|
||||
Tags && ...tags)
|
||||
{
|
||||
using input_type = std::decay_t<InputT>;
|
||||
detail::ic_impl::check_input<input_type>();
|
||||
|
|
@ -536,7 +554,8 @@ auto make_dpf_doerner_shelat(InputT r0, InputT r1,
|
|||
const input_type g0 = r0;
|
||||
const input_type g1 = utils::xor_input_shares(g0, gamma);
|
||||
auto inner = make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(g0, g1,
|
||||
std::move(rng), detail::ic_impl::inner_lt(spec));
|
||||
std::move(rng), detail::ic_impl::inner_lt(spec),
|
||||
std::forward<Tags>(tags)...);
|
||||
return detail::ic_impl::finish<input_type>(r_bits, spec, std::move(inner));
|
||||
}
|
||||
|
||||
|
|
@ -547,16 +566,23 @@ auto make_dpf_doerner_shelat(InputT r0, InputT r1,
|
|||
/// @param r1 party 1's share of the mask
|
||||
/// @param rng the Doerner–Shelat randomness tapes
|
||||
/// @param spec the public bounds and payloads
|
||||
/// @param tags optional `verifiable` or `extractable` markers
|
||||
/// @return the two party keys
|
||||
/// \complexity O(n) time. One `ds_advance_level` per level: two PRG expansions and one `prepare_level`. n is `depth`.
|
||||
/// \rounds No sockets. This is the in-process transcript. A networked walk is `dpf::party::dist::point_party`.
|
||||
/// \communication none here. `local_cw_protocol` opens the correction word locally.
|
||||
/// \preprocessing Per level, `prepare_level` draws one `ds_cw_pads` (two parties × a 128-bit rand, a 128-bit gamma, and a bit) and two `ds_and_pads`. Arithmetic inputs also run a carry chain of n-1 bit-AND triples in `encode_walk_shares`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename RootSampler,
|
||||
typename PadRng,
|
||||
typename InputT,
|
||||
typename Beta>
|
||||
typename Beta,
|
||||
typename ...Tags>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto make_dpf_doerner_shelat(arith_input_t, InputT r0, InputT r1,
|
||||
ds_randomness<RootSampler, PadRng> rng, const ic_pack<Beta> & spec)
|
||||
ds_randomness<RootSampler, PadRng> rng, const ic_pack<Beta> & spec,
|
||||
Tags && ...tags)
|
||||
{
|
||||
using input_type = std::decay_t<InputT>;
|
||||
detail::ic_impl::check_input<input_type>();
|
||||
|
|
@ -570,42 +596,56 @@ auto make_dpf_doerner_shelat(arith_input_t, InputT r0, InputT r1,
|
|||
(detail::ic_impl::bits_of(r0) - 1ULL) & nmask);
|
||||
const input_type g1 = r1;
|
||||
auto inner = make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
|
||||
arith_input, g0, g1, std::move(rng), detail::ic_impl::inner_lt(spec));
|
||||
arith_input, g0, g1, std::move(rng), detail::ic_impl::inner_lt(spec),
|
||||
std::forward<Tags>(tags)...);
|
||||
return detail::ic_impl::finish<input_type>(
|
||||
detail::ic_impl::bits_of(r), spec, std::move(inner));
|
||||
}
|
||||
|
||||
/// @brief XOR shares, sampled from the library entropy source.
|
||||
/// @return the two party keys
|
||||
/// \complexity O(n) time. One `ds_advance_level` per level: two PRG expansions and one `prepare_level`. n is `depth`.
|
||||
/// \rounds No sockets. This is the in-process transcript. A networked walk is `dpf::party::dist::point_party`.
|
||||
/// \communication none here. `local_cw_protocol` opens the correction word locally.
|
||||
/// \preprocessing Per level, `prepare_level` draws one `ds_cw_pads` (two parties × a 128-bit rand, a 128-bit gamma, and a bit) and two `ds_and_pads`. Arithmetic inputs also run a carry chain of n-1 bit-AND triples in `encode_walk_shares`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
typename Beta>
|
||||
typename Beta,
|
||||
typename ...Tags>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto make_dpf_doerner_shelat(InputT r0, InputT r1, const ic_pack<Beta> & spec)
|
||||
auto make_dpf_doerner_shelat(InputT r0, InputT r1, const ic_pack<Beta> & spec,
|
||||
Tags && ...tags)
|
||||
{
|
||||
using block = typename InteriorPRG::block_type;
|
||||
ds_randomness<block (*)(), detail::urandom_pad_rng> rng{
|
||||
dpf::uniform_sample<block>, {}};
|
||||
return make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
|
||||
std::move(r0), std::move(r1), rng, spec);
|
||||
std::move(r0), std::move(r1), rng, spec,
|
||||
std::forward<Tags>(tags)...);
|
||||
}
|
||||
|
||||
/// @brief Additive shares, sampled from the library entropy source.
|
||||
/// @return the two party keys
|
||||
/// \complexity O(n) time. One `ds_advance_level` per level: two PRG expansions and one `prepare_level`. n is `depth`.
|
||||
/// \rounds No sockets. This is the in-process transcript. A networked walk is `dpf::party::dist::point_party`.
|
||||
/// \communication none here. `local_cw_protocol` opens the correction word locally.
|
||||
/// \preprocessing Per level, `prepare_level` draws one `ds_cw_pads` (two parties × a 128-bit rand, a 128-bit gamma, and a bit) and two `ds_and_pads`. Arithmetic inputs also run a carry chain of n-1 bit-AND triples in `encode_walk_shares`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename InputT,
|
||||
typename Beta>
|
||||
typename Beta,
|
||||
typename ...Tags>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto make_dpf_doerner_shelat(arith_input_t, InputT r0, InputT r1,
|
||||
const ic_pack<Beta> & spec)
|
||||
const ic_pack<Beta> & spec, Tags && ...tags)
|
||||
{
|
||||
using block = typename InteriorPRG::block_type;
|
||||
ds_randomness<block (*)(), detail::urandom_pad_rng> rng{
|
||||
dpf::uniform_sample<block>, {}};
|
||||
return make_dpf_doerner_shelat<InteriorPRG, ExteriorPRG>(
|
||||
arith_input, std::move(r0), std::move(r1), rng, spec);
|
||||
arith_input, std::move(r0), std::move(r1), rng, spec,
|
||||
std::forward<Tags>(tags)...);
|
||||
}
|
||||
|
||||
/// @}
|
||||
|
|
@ -633,7 +673,7 @@ void assign_cmp(ic_key<0, Key, Input, Beta> & k0,
|
|||
const auto delta = detail::group_sub(
|
||||
detail::group_from_beta(if_true), detail::group_from_beta(if_false));
|
||||
const auto fval = detail::group_from_beta(if_false);
|
||||
assign_cmp(k0.key, k1.key,
|
||||
assign_cmp(k0.dpf_key, k1.dpf_key,
|
||||
detail::group_to_beta<Beta>(delta),
|
||||
detail::group_to_beta<Beta>(detail::group_zero(layout)));
|
||||
k0.delta_share = detail::group_mul(k0.delta_coeff, delta);
|
||||
|
|
@ -656,7 +696,7 @@ void assign_cmp(ic_key<0, Key, Input, Beta> & k0,
|
|||
detail::dcf_impl::beta_delta_u64(if_true, if_false, mask);
|
||||
const uint64_t fval =
|
||||
detail::dcf_impl::beta_to_u64_simple(if_false, mask);
|
||||
assign_cmp(k0.key, k1.key,
|
||||
assign_cmp(k0.dpf_key, k1.dpf_key,
|
||||
detail::dcf_impl::u64_to_beta<Beta>(delta),
|
||||
detail::dcf_impl::u64_to_beta<Beta>(0));
|
||||
k0.delta_share = detail::ic_impl::mul_mask(k0.delta_coeff, delta, mask);
|
||||
|
|
@ -680,6 +720,7 @@ void assign_cmp(ic_key<0, Key, Input, Beta> & k0,
|
|||
/// @param x the `x`
|
||||
/// @param memo the memoizer reused across queries
|
||||
/// @return Point evaluation
|
||||
/// \complexity Two comparison point-walks (`eval_cmp_point_impl`), each O(n) interior steps, plus O(1) group arithmetic. n is the key depth. A path memoizer reuses a shared prefix.
|
||||
template <typename IcKey, typename Query,
|
||||
typename Memo = basic_path_memoizer<typename IcKey::key_type>,
|
||||
typename = std::enable_if_t<is_ic_key_v<IcKey>>>
|
||||
|
|
@ -707,6 +748,7 @@ auto eval_point(ic_fn, const IcKey & key, Query && x, Memo && memo = Memo{})
|
|||
/// @param to the inclusive end of the range
|
||||
/// @param buf the output buffer
|
||||
/// @param memo the memoizer reused across queries
|
||||
/// \complexity One `eval_one` per input from `from` through `to`. Each `eval_one` is two comparison point-walks, so this is not the truncated-tree interval walk. A path memoizer reuses prefixes across those walks. Counted the `for (x = a; x != b; ++x)` loop.
|
||||
template <typename IcKey, typename Lane, typename Buffer, typename Memo,
|
||||
typename = std::enable_if_t<is_ic_key_v<IcKey>>>
|
||||
void eval_interval(ic_fn, const IcKey & key, Lane from, Lane to,
|
||||
|
|
@ -731,6 +773,7 @@ void eval_interval(ic_fn, const IcKey & key, Lane from, Lane to,
|
|||
}
|
||||
|
||||
/// @brief Inclusive interval `[from, to]`, with a fresh path memoizer.
|
||||
/// \complexity One `eval_one` per input from `from` through `to`. Each `eval_one` is two comparison point-walks, so this is not the truncated-tree interval walk. A path memoizer reuses prefixes across those walks. Counted the `for (x = a; x != b; ++x)` loop.
|
||||
template <typename IcKey, typename Lane, typename Buffer,
|
||||
typename = std::enable_if_t<is_ic_key_v<IcKey>>>
|
||||
void eval_interval(ic_fn, const IcKey & key, Lane from, Lane to, Buffer && buf)
|
||||
|
|
@ -758,6 +801,7 @@ void eval_interval(ic_fn, const IcKey & key, Lane from, Lane to, Buffer && buf)
|
|||
/// @param end the iterator past the last query
|
||||
/// @param buf the output buffer
|
||||
/// @param memo the memoizer reused across queries
|
||||
/// \complexity One `eval_one` per input from `from` through `to`. Each `eval_one` is two comparison point-walks, so this is not the truncated-tree interval walk. A path memoizer reuses prefixes across those walks. Counted the `for (x = a; x != b; ++x)` loop.
|
||||
template <typename IcKey, typename Iter, typename Buffer, typename Memo,
|
||||
typename = std::enable_if_t<is_ic_key_v<IcKey>>>
|
||||
void eval_sequence(ic_fn, const IcKey & key, Iter begin, Iter end,
|
||||
|
|
@ -769,6 +813,7 @@ void eval_sequence(ic_fn, const IcKey & key, Iter begin, Iter end,
|
|||
}
|
||||
|
||||
/// @brief Evaluate `[begin, end)`, with a fresh path memoizer.
|
||||
/// \complexity One `eval_one` per input from `from` through `to`. Each `eval_one` is two comparison point-walks, so this is not the truncated-tree interval walk. A path memoizer reuses prefixes across those walks. Counted the `for (x = a; x != b; ++x)` loop.
|
||||
template <typename IcKey, typename Iter, typename Buffer,
|
||||
typename = std::enable_if_t<is_ic_key_v<IcKey>>>
|
||||
void eval_sequence(ic_fn, const IcKey & key, Iter begin, Iter end, Buffer && buf)
|
||||
|
|
@ -829,8 +874,70 @@ auto make_output_buffer(ic_fn, const IcKey & key, Lane from, Lane to)
|
|||
/// @return the opened party shares
|
||||
/// @{
|
||||
|
||||
namespace detail
|
||||
{
|
||||
namespace ic_impl
|
||||
{
|
||||
|
||||
/// @brief Fill `out` from a verifiable IC key pair, reusing one path memoizer.
|
||||
template <typename IcKey0, typename IcKey1, typename Iter>
|
||||
void geneval_ic_eval(geneval_cmp_result & out, const IcKey0 & k0,
|
||||
const IcKey1 & k1, Iter begin, Iter end)
|
||||
{
|
||||
using key_type = unwrap_party_key_t<typename IcKey0::key_type>;
|
||||
constexpr std::size_t depth = key_type::depth;
|
||||
out.live_levels = depth;
|
||||
out.mask = k0.dpf_key.cmp().mask;
|
||||
out.cw_last = k0.dpf_key.cw_last();
|
||||
out.addend0 = k0.dpf_key.cmp_addend().raw();
|
||||
out.addend1 = k1.dpf_key.cmp_addend().raw();
|
||||
out.correction_words.resize(depth);
|
||||
out.correction_advice.resize(depth);
|
||||
if constexpr (key_type::cmp_block > 0)
|
||||
{
|
||||
out.value_cw.resize(key_type::cmp_checkpoints);
|
||||
for (std::size_t i = 0; i < key_type::cmp_checkpoints; ++i)
|
||||
out.value_cw[i] = k0.dpf_key.value_cw(i);
|
||||
out.tail_cw.resize(key_type::cmp_tail);
|
||||
for (std::size_t z = 0; z < key_type::cmp_tail; ++z)
|
||||
out.tail_cw[z] = k0.dpf_key.tail_cw(z);
|
||||
}
|
||||
else
|
||||
out.value_cw.resize(depth);
|
||||
for (std::size_t level = 0; level < depth; ++level)
|
||||
{
|
||||
out.correction_words[level] = k0.dpf_key.correction_word(level);
|
||||
out.correction_advice[level] =
|
||||
static_cast<uint8_t>(k0.dpf_key.correction_advice(level));
|
||||
if constexpr (key_type::cmp_block == 0)
|
||||
out.value_cw[level] = k0.dpf_key.value_cw(level);
|
||||
}
|
||||
// Empty query lists keep default-constructed (zero) tokens. `verify`
|
||||
// rejects the all-zero token, so an empty geneval does not verify as
|
||||
// two matching zeros.
|
||||
detail::vdpf::init_proof(out.proof0, k0.dpf_key);
|
||||
detail::vdpf::init_proof(out.proof1, k1.dpf_key);
|
||||
auto path0 = make_basic_path_memoizer(k0.dpf_key);
|
||||
auto path1 = make_basic_path_memoizer(k1.dpf_key);
|
||||
for (auto it = begin; it != end; ++it)
|
||||
{
|
||||
out.party0.push_back(opened_u64(
|
||||
eval_one(k0, *it, path0, &out.proof0), out.mask));
|
||||
out.party1.push_back(opened_u64(
|
||||
eval_one(k1, *it, path1, &out.proof1), out.mask));
|
||||
}
|
||||
detail::vdpf::fold_output_binding(out.proof0, k0.dpf_key);
|
||||
detail::vdpf::fold_output_binding(out.proof1, k1.dpf_key);
|
||||
}
|
||||
|
||||
} // namespace ic_impl
|
||||
} // namespace detail
|
||||
|
||||
/// @brief XOR mask. `r0 XOR r1` is the secret mask. Each query is
|
||||
/// returned already combined into the interval share.
|
||||
/// @details Builds a verifiable inner comparison key and folds correction
|
||||
/// seeds into `proof0` / `proof1` while reusing one path memoizer per party.
|
||||
/// An empty query range leaves both tokens zero; those do not verify.
|
||||
template <typename InputT, typename Iter, typename RootSampler, typename PadRng,
|
||||
typename Beta>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
|
|
@ -844,33 +951,8 @@ geneval_cmp_result geneval_ic(InputT r0, InputT r1, Iter begin, Iter end,
|
|||
return out;
|
||||
|
||||
auto keys = make_dpf_doerner_shelat(std::move(r0), std::move(r1),
|
||||
std::move(rng), spec);
|
||||
const auto & k0 = keys.first;
|
||||
const auto & k1 = keys.second;
|
||||
using key_type = unwrap_party_key_t<typename std::decay_t<decltype(k0)>::key_type>;
|
||||
constexpr std::size_t depth = key_type::depth;
|
||||
out.live_levels = depth;
|
||||
out.mask = k0.key.cmp().mask;
|
||||
out.cw_last = k0.key.cw_last();
|
||||
out.addend0 = k0.key.cmp_addend().raw();
|
||||
out.addend1 = k1.key.cmp_addend().raw();
|
||||
out.correction_words.resize(depth);
|
||||
out.correction_advice.resize(depth);
|
||||
out.value_cw.resize(depth);
|
||||
for (std::size_t level = 0; level < depth; ++level)
|
||||
{
|
||||
out.correction_words[level] = k0.key.correction_word(level);
|
||||
out.correction_advice[level] =
|
||||
static_cast<uint8_t>(k0.key.correction_advice(level));
|
||||
out.value_cw[level] = k0.key.value_cw(level);
|
||||
}
|
||||
for (auto it = begin; it != end; ++it)
|
||||
{
|
||||
out.party0.push_back(detail::ic_impl::opened_u64(
|
||||
eval_point(ic, k0, *it), out.mask));
|
||||
out.party1.push_back(detail::ic_impl::opened_u64(
|
||||
eval_point(ic, k1, *it), out.mask));
|
||||
}
|
||||
std::move(rng), spec, dpf::verifiable{});
|
||||
detail::ic_impl::geneval_ic_eval(out, keys.first, keys.second, begin, end);
|
||||
return out;
|
||||
}
|
||||
|
||||
|
|
@ -888,38 +970,339 @@ geneval_cmp_result geneval_ic(arith_input_t, InputT r0, InputT r1, Iter begin,
|
|||
return out;
|
||||
|
||||
auto keys = make_dpf_doerner_shelat(arith_input, std::move(r0), std::move(r1),
|
||||
std::move(rng), spec);
|
||||
const auto & k0 = keys.first;
|
||||
const auto & k1 = keys.second;
|
||||
using key_type = unwrap_party_key_t<typename std::decay_t<decltype(k0)>::key_type>;
|
||||
constexpr std::size_t depth = key_type::depth;
|
||||
out.live_levels = depth;
|
||||
out.mask = k0.key.cmp().mask;
|
||||
out.cw_last = k0.key.cw_last();
|
||||
out.addend0 = k0.key.cmp_addend().raw();
|
||||
out.addend1 = k1.key.cmp_addend().raw();
|
||||
out.correction_words.resize(depth);
|
||||
out.correction_advice.resize(depth);
|
||||
out.value_cw.resize(depth);
|
||||
for (std::size_t level = 0; level < depth; ++level)
|
||||
{
|
||||
out.correction_words[level] = k0.key.correction_word(level);
|
||||
out.correction_advice[level] =
|
||||
static_cast<uint8_t>(k0.key.correction_advice(level));
|
||||
out.value_cw[level] = k0.key.value_cw(level);
|
||||
}
|
||||
for (auto it = begin; it != end; ++it)
|
||||
{
|
||||
out.party0.push_back(detail::ic_impl::opened_u64(
|
||||
eval_point(ic, k0, *it), out.mask));
|
||||
out.party1.push_back(detail::ic_impl::opened_u64(
|
||||
eval_point(ic, k1, *it), out.mask));
|
||||
}
|
||||
std::move(rng), spec, dpf::verifiable{});
|
||||
detail::ic_impl::geneval_ic_eval(out, keys.first, keys.second, begin, end);
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Grow adaptations: run on the inner `dpf_key`, then refresh `cr_share` from
|
||||
// the secret mask `r` when depth changes (`add_output` copies the fields).
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
template <std::size_t P0, std::size_t P1, typename Key, typename Input,
|
||||
typename Beta, typename InputT, typename... Specs>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto extend(const ic_key<P0, Key, Input, Beta> & k0,
|
||||
const ic_key<P1, Key, Input, Beta> & k1, uint64_t r, bool bit, InputT x,
|
||||
Specs &&... specs)
|
||||
{
|
||||
auto inner = dpf::extend(k0.dpf_key, k1.dpf_key, bit, x,
|
||||
std::forward<Specs>(specs)...);
|
||||
using new_raw = typename std::decay_t<decltype(inner.first)>::key_type;
|
||||
const uint64_t nmask = detail::ic_impl::input_mask_of<Input>();
|
||||
const uint64_t gmask = k0.group_mask;
|
||||
using share_t = typename ic_key<P0, new_raw, Input, Beta>::share_type;
|
||||
share_t c0{};
|
||||
share_t c1{};
|
||||
if constexpr (is_wildcard_v<Beta>)
|
||||
{
|
||||
// Wildcard: keep cr_coeff; concrete cr_share is filled by assign_cmp.
|
||||
c0 = k0.cr_share;
|
||||
c1 = k1.cr_share;
|
||||
}
|
||||
else if constexpr (detail::cmp_group_info<Beta>::custom)
|
||||
{
|
||||
using prg = typename new_raw::interior_prg;
|
||||
const auto layout = detail::group_layout<concrete_type_t<Beta>>();
|
||||
const auto cr_g = detail::group_scalar(
|
||||
detail::ic_impl::correction_s(r, k0.lo, k0.hi, nmask), layout);
|
||||
const auto target = detail::group_add(
|
||||
detail::group_mul(k0.delta_share, cr_g), // wrong: need open δ
|
||||
detail::group_zero(layout));
|
||||
(void)target;
|
||||
// Reconstruct δ = d0+d1, if_false from old cr, then new cr = δ·c_r + f.
|
||||
const auto delta = detail::group_add(k0.delta_share, k1.delta_share);
|
||||
const auto old_cr_open = detail::group_add(k0.cr_share, k1.cr_share);
|
||||
const auto old_cr_term = detail::group_mul(delta,
|
||||
detail::group_scalar(
|
||||
detail::ic_impl::correction_s(r, k0.lo, k0.hi, k0.input_mask),
|
||||
layout));
|
||||
const auto if_false = detail::group_sub(old_cr_open, old_cr_term);
|
||||
const auto new_target =
|
||||
detail::group_add(detail::group_mul(delta, cr_g), if_false);
|
||||
const auto blind = detail::group_from_node<prg>(
|
||||
dpf::uniform_sample<typename new_raw::interior_node>(), layout);
|
||||
c0 = blind;
|
||||
c1 = detail::group_sub(new_target, blind);
|
||||
}
|
||||
else
|
||||
{
|
||||
const uint64_t delta =
|
||||
(static_cast<uint64_t>(k0.delta_share)
|
||||
+ static_cast<uint64_t>(k1.delta_share))
|
||||
& gmask;
|
||||
const uint64_t old_cr_open =
|
||||
(static_cast<uint64_t>(k0.cr_share)
|
||||
+ static_cast<uint64_t>(k1.cr_share))
|
||||
& gmask;
|
||||
const uint64_t old_cr_term = detail::ic_impl::mul_mask(delta,
|
||||
detail::ic_impl::correction(r, k0.lo, k0.hi, k0.input_mask, gmask),
|
||||
gmask);
|
||||
const uint64_t if_false =
|
||||
(old_cr_open + detail::dcf_impl::neg_m(old_cr_term, gmask)) & gmask;
|
||||
const uint64_t cr =
|
||||
detail::ic_impl::correction(r, k0.lo, k0.hi, nmask, gmask);
|
||||
const uint64_t target =
|
||||
(detail::ic_impl::mul_mask(delta, cr, gmask) + if_false) & gmask;
|
||||
const uint64_t blind = detail::dcf_impl::sample_addend_blind(gmask,
|
||||
[] {
|
||||
return dpf::uniform_sample<typename new_raw::interior_node>();
|
||||
});
|
||||
uint64_t a0 = 0, a1 = 0;
|
||||
detail::incr::split_cmp_addend(target, gmask, blind, a0, a1);
|
||||
c0 = static_cast<share_t>(a0);
|
||||
c1 = static_cast<share_t>(a1);
|
||||
}
|
||||
auto out0 = detail::ic_impl::make_side<P0, new_raw, Input, Beta>(
|
||||
std::move(inner.first), k0.lo, k0.hi, nmask, gmask, k0.delta_share, c0,
|
||||
k0.delta_coeff, k0.cr_coeff);
|
||||
auto out1 = detail::ic_impl::make_side<P1, new_raw, Input, Beta>(
|
||||
std::move(inner.second), k1.lo, k1.hi, nmask, gmask, k1.delta_share, c1,
|
||||
k1.delta_coeff, k1.cr_coeff);
|
||||
out0.assigned = k0.assigned;
|
||||
out1.assigned = k1.assigned;
|
||||
return std::make_pair(std::move(out0), std::move(out1));
|
||||
}
|
||||
|
||||
template <std::size_t P0, std::size_t P1, typename Key, typename Input,
|
||||
typename Beta, typename InputT, typename... Specs>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto extend(const ic_key<P0, Key, Input, Beta> & k0,
|
||||
const ic_key<P1, Key, Input, Beta> & k1, uint64_t r, InputT x,
|
||||
Specs &&... specs)
|
||||
{
|
||||
const bool bit = detail::grow_impl::bit_at(
|
||||
static_cast<typename Key::input_type>(x), Key::depth);
|
||||
return extend(k0, k1, r, bit, x, std::forward<Specs>(specs)...);
|
||||
}
|
||||
|
||||
template <std::size_t P0, std::size_t P1, typename Key, typename Input,
|
||||
typename Beta, typename InputT, typename... Specs>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto add_output(const ic_key<P0, Key, Input, Beta> & k0,
|
||||
const ic_key<P1, Key, Input, Beta> & k1, InputT x, Specs &&... specs)
|
||||
{
|
||||
auto inner = dpf::add_output(k0.dpf_key, k1.dpf_key, x,
|
||||
std::forward<Specs>(specs)...);
|
||||
using new_raw = typename std::decay_t<decltype(inner.first)>::key_type;
|
||||
auto out0 = detail::ic_impl::make_side<P0, new_raw, Input, Beta>(
|
||||
std::move(inner.first), k0.lo, k0.hi, k0.input_mask, k0.group_mask,
|
||||
k0.delta_share, k0.cr_share, k0.delta_coeff, k0.cr_coeff);
|
||||
auto out1 = detail::ic_impl::make_side<P1, new_raw, Input, Beta>(
|
||||
std::move(inner.second), k1.lo, k1.hi, k1.input_mask, k1.group_mask,
|
||||
k1.delta_share, k1.cr_share, k1.delta_coeff, k1.cr_coeff);
|
||||
out0.assigned = k0.assigned;
|
||||
out1.assigned = k1.assigned;
|
||||
return std::make_pair(std::move(out0), std::move(out1));
|
||||
}
|
||||
|
||||
template <std::size_t P0, std::size_t P1, typename Key, typename Input,
|
||||
typename Beta, typename Memo0, typename Memo1, typename InputT,
|
||||
typename... Specs>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto extend(const ic_key<P0, Key, Input, Beta> & k0,
|
||||
const ic_key<P1, Key, Input, Beta> & k1, uint64_t r, Memo0 & m0, Memo1 & m1,
|
||||
bool bit, InputT x, Specs &&... specs)
|
||||
{
|
||||
auto inner = dpf::extend(k0.dpf_key, k1.dpf_key, m0, m1, bit, x,
|
||||
std::forward<Specs>(specs)...);
|
||||
using new_raw = typename std::decay_t<decltype(inner.first)>::key_type;
|
||||
const uint64_t nmask = detail::ic_impl::input_mask_of<Input>();
|
||||
const uint64_t gmask = k0.group_mask;
|
||||
using share_t = typename ic_key<P0, new_raw, Input, Beta>::share_type;
|
||||
share_t c0{};
|
||||
share_t c1{};
|
||||
if constexpr (is_wildcard_v<Beta>)
|
||||
{
|
||||
c0 = k0.cr_share;
|
||||
c1 = k1.cr_share;
|
||||
}
|
||||
else if constexpr (detail::cmp_group_info<Beta>::custom)
|
||||
{
|
||||
using prg = typename new_raw::interior_prg;
|
||||
const auto layout = detail::group_layout<concrete_type_t<Beta>>();
|
||||
const auto cr_g = detail::group_scalar(
|
||||
detail::ic_impl::correction_s(r, k0.lo, k0.hi, nmask), layout);
|
||||
const auto delta = detail::group_add(k0.delta_share, k1.delta_share);
|
||||
const auto old_cr_open = detail::group_add(k0.cr_share, k1.cr_share);
|
||||
const auto old_cr_term = detail::group_mul(delta,
|
||||
detail::group_scalar(
|
||||
detail::ic_impl::correction_s(r, k0.lo, k0.hi, k0.input_mask),
|
||||
layout));
|
||||
const auto if_false = detail::group_sub(old_cr_open, old_cr_term);
|
||||
const auto new_target =
|
||||
detail::group_add(detail::group_mul(delta, cr_g), if_false);
|
||||
const auto blind = detail::group_from_node<prg>(
|
||||
dpf::uniform_sample<typename new_raw::interior_node>(), layout);
|
||||
c0 = blind;
|
||||
c1 = detail::group_sub(new_target, blind);
|
||||
}
|
||||
else
|
||||
{
|
||||
const uint64_t delta =
|
||||
(static_cast<uint64_t>(k0.delta_share)
|
||||
+ static_cast<uint64_t>(k1.delta_share))
|
||||
& gmask;
|
||||
const uint64_t old_cr_open =
|
||||
(static_cast<uint64_t>(k0.cr_share)
|
||||
+ static_cast<uint64_t>(k1.cr_share))
|
||||
& gmask;
|
||||
const uint64_t old_cr_term = detail::ic_impl::mul_mask(delta,
|
||||
detail::ic_impl::correction(r, k0.lo, k0.hi, k0.input_mask, gmask),
|
||||
gmask);
|
||||
const uint64_t if_false =
|
||||
(old_cr_open + detail::dcf_impl::neg_m(old_cr_term, gmask)) & gmask;
|
||||
const uint64_t cr =
|
||||
detail::ic_impl::correction(r, k0.lo, k0.hi, nmask, gmask);
|
||||
const uint64_t target =
|
||||
(detail::ic_impl::mul_mask(delta, cr, gmask) + if_false) & gmask;
|
||||
const uint64_t blind = detail::dcf_impl::sample_addend_blind(gmask,
|
||||
[] {
|
||||
return dpf::uniform_sample<typename new_raw::interior_node>();
|
||||
});
|
||||
uint64_t a0 = 0, a1 = 0;
|
||||
detail::incr::split_cmp_addend(target, gmask, blind, a0, a1);
|
||||
c0 = static_cast<share_t>(a0);
|
||||
c1 = static_cast<share_t>(a1);
|
||||
}
|
||||
auto out0 = detail::ic_impl::make_side<P0, new_raw, Input, Beta>(
|
||||
std::move(inner.first), k0.lo, k0.hi, nmask, gmask, k0.delta_share, c0,
|
||||
k0.delta_coeff, k0.cr_coeff);
|
||||
auto out1 = detail::ic_impl::make_side<P1, new_raw, Input, Beta>(
|
||||
std::move(inner.second), k1.lo, k1.hi, nmask, gmask, k1.delta_share, c1,
|
||||
k1.delta_coeff, k1.cr_coeff);
|
||||
out0.assigned = k0.assigned;
|
||||
out1.assigned = k1.assigned;
|
||||
return std::make_pair(std::move(out0), std::move(out1));
|
||||
}
|
||||
|
||||
template <std::size_t P0, std::size_t P1, typename Key, typename Input,
|
||||
typename Beta, typename Memo0, typename Memo1, typename InputT,
|
||||
typename... Specs>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto add_output(const ic_key<P0, Key, Input, Beta> & k0,
|
||||
const ic_key<P1, Key, Input, Beta> & k1, Memo0 & m0, Memo1 & m1, InputT x,
|
||||
Specs &&... specs)
|
||||
{
|
||||
auto inner = dpf::add_output(k0.dpf_key, k1.dpf_key, m0, m1, x,
|
||||
std::forward<Specs>(specs)...);
|
||||
using new_raw = typename std::decay_t<decltype(inner.first)>::key_type;
|
||||
auto out0 = detail::ic_impl::make_side<P0, new_raw, Input, Beta>(
|
||||
std::move(inner.first), k0.lo, k0.hi, k0.input_mask, k0.group_mask,
|
||||
k0.delta_share, k0.cr_share, k0.delta_coeff, k0.cr_coeff);
|
||||
auto out1 = detail::ic_impl::make_side<P1, new_raw, Input, Beta>(
|
||||
std::move(inner.second), k1.lo, k1.hi, k1.input_mask, k1.group_mask,
|
||||
k1.delta_share, k1.cr_share, k1.delta_coeff, k1.cr_coeff);
|
||||
out0.assigned = k0.assigned;
|
||||
out1.assigned = k1.assigned;
|
||||
return std::make_pair(std::move(out0), std::move(out1));
|
||||
}
|
||||
|
||||
template <std::size_t P0, std::size_t P1, typename Key, typename Input,
|
||||
typename Beta, typename Memo0, typename Memo1, typename InputT,
|
||||
typename CwProtocol, typename... Specs>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto extend_ds(const ic_key<P0, Key, Input, Beta> & k0,
|
||||
const ic_key<P1, Key, Input, Beta> & k1, uint64_t r, Memo0 & m0, Memo1 & m1,
|
||||
InputT x0, InputT x1, CwProtocol & proto, Specs &&... specs)
|
||||
{
|
||||
auto inner = dpf::extend_ds(k0.dpf_key, k1.dpf_key, m0, m1, x0, x1, proto,
|
||||
std::forward<Specs>(specs)...);
|
||||
using new_raw = typename std::decay_t<decltype(inner.first)>::key_type;
|
||||
const uint64_t nmask = detail::ic_impl::input_mask_of<Input>();
|
||||
const uint64_t gmask = k0.group_mask;
|
||||
using share_t = typename ic_key<P0, new_raw, Input, Beta>::share_type;
|
||||
share_t c0{};
|
||||
share_t c1{};
|
||||
if constexpr (is_wildcard_v<Beta>)
|
||||
{
|
||||
c0 = k0.cr_share;
|
||||
c1 = k1.cr_share;
|
||||
}
|
||||
else if constexpr (detail::cmp_group_info<Beta>::custom)
|
||||
{
|
||||
using prg = typename new_raw::interior_prg;
|
||||
const auto layout = detail::group_layout<concrete_type_t<Beta>>();
|
||||
const auto cr_g = detail::group_scalar(
|
||||
detail::ic_impl::correction_s(r, k0.lo, k0.hi, nmask), layout);
|
||||
const auto delta = detail::group_add(k0.delta_share, k1.delta_share);
|
||||
const auto old_cr_open = detail::group_add(k0.cr_share, k1.cr_share);
|
||||
const auto old_cr_term = detail::group_mul(delta,
|
||||
detail::group_scalar(
|
||||
detail::ic_impl::correction_s(r, k0.lo, k0.hi, k0.input_mask),
|
||||
layout));
|
||||
const auto if_false = detail::group_sub(old_cr_open, old_cr_term);
|
||||
const auto new_target =
|
||||
detail::group_add(detail::group_mul(delta, cr_g), if_false);
|
||||
const auto blind = detail::group_from_node<prg>(
|
||||
dpf::uniform_sample<typename new_raw::interior_node>(), layout);
|
||||
c0 = blind;
|
||||
c1 = detail::group_sub(new_target, blind);
|
||||
}
|
||||
else
|
||||
{
|
||||
const uint64_t delta =
|
||||
(static_cast<uint64_t>(k0.delta_share)
|
||||
+ static_cast<uint64_t>(k1.delta_share))
|
||||
& gmask;
|
||||
const uint64_t old_cr_open =
|
||||
(static_cast<uint64_t>(k0.cr_share)
|
||||
+ static_cast<uint64_t>(k1.cr_share))
|
||||
& gmask;
|
||||
const uint64_t old_cr_term = detail::ic_impl::mul_mask(delta,
|
||||
detail::ic_impl::correction(r, k0.lo, k0.hi, k0.input_mask, gmask),
|
||||
gmask);
|
||||
const uint64_t if_false =
|
||||
(old_cr_open + detail::dcf_impl::neg_m(old_cr_term, gmask)) & gmask;
|
||||
const uint64_t cr =
|
||||
detail::ic_impl::correction(r, k0.lo, k0.hi, nmask, gmask);
|
||||
const uint64_t target =
|
||||
(detail::ic_impl::mul_mask(delta, cr, gmask) + if_false) & gmask;
|
||||
const uint64_t blind = detail::dcf_impl::sample_addend_blind(gmask,
|
||||
[] {
|
||||
return dpf::uniform_sample<typename new_raw::interior_node>();
|
||||
});
|
||||
uint64_t a0 = 0, a1 = 0;
|
||||
detail::incr::split_cmp_addend(target, gmask, blind, a0, a1);
|
||||
c0 = static_cast<share_t>(a0);
|
||||
c1 = static_cast<share_t>(a1);
|
||||
}
|
||||
auto out0 = detail::ic_impl::make_side<P0, new_raw, Input, Beta>(
|
||||
std::move(inner.first), k0.lo, k0.hi, nmask, gmask, k0.delta_share, c0,
|
||||
k0.delta_coeff, k0.cr_coeff);
|
||||
auto out1 = detail::ic_impl::make_side<P1, new_raw, Input, Beta>(
|
||||
std::move(inner.second), k1.lo, k1.hi, nmask, gmask, k1.delta_share, c1,
|
||||
k1.delta_coeff, k1.cr_coeff);
|
||||
out0.assigned = k0.assigned;
|
||||
out1.assigned = k1.assigned;
|
||||
return std::make_pair(std::move(out0), std::move(out1));
|
||||
}
|
||||
|
||||
template <std::size_t P0, std::size_t P1, typename Key, typename Input,
|
||||
typename Beta, typename Memo0, typename Memo1, typename InputT,
|
||||
typename CwProtocol, typename... Specs>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto add_output_ds(const ic_key<P0, Key, Input, Beta> & k0,
|
||||
const ic_key<P1, Key, Input, Beta> & k1, Memo0 & m0, Memo1 & m1, InputT x0,
|
||||
InputT x1, CwProtocol & proto, Specs &&... specs)
|
||||
{
|
||||
auto inner = dpf::add_output_ds(k0.dpf_key, k1.dpf_key, m0, m1, x0, x1, proto,
|
||||
std::forward<Specs>(specs)...);
|
||||
using new_raw = typename std::decay_t<decltype(inner.first)>::key_type;
|
||||
auto out0 = detail::ic_impl::make_side<P0, new_raw, Input, Beta>(
|
||||
std::move(inner.first), k0.lo, k0.hi, k0.input_mask, k0.group_mask,
|
||||
k0.delta_share, k0.cr_share, k0.delta_coeff, k0.cr_coeff);
|
||||
auto out1 = detail::ic_impl::make_side<P1, new_raw, Input, Beta>(
|
||||
std::move(inner.second), k1.lo, k1.hi, k1.input_mask, k1.group_mask,
|
||||
k1.delta_share, k1.cr_share, k1.delta_coeff, k1.cr_coeff);
|
||||
out0.assigned = k0.assigned;
|
||||
out1.assigned = k1.assigned;
|
||||
return std::make_pair(std::move(out0), std::move(out1));
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_INTERVAL_HPP__
|
||||
|
|
|
|||
|
|
@ -73,6 +73,16 @@ struct interval_memoizer_base
|
|||
HEDLEY_NO_THROW
|
||||
virtual return_type end() const noexcept = 0;
|
||||
|
||||
/// @brief Drop a cached interval so the next `assign_interval` rebuilds
|
||||
/// from the root (needed when folding a proof over the tree).
|
||||
void clear_assignment()
|
||||
{
|
||||
dpf_ = std::nullopt;
|
||||
from_ = std::nullopt;
|
||||
to_ = std::nullopt;
|
||||
level_index = 0;
|
||||
}
|
||||
|
||||
virtual std::size_t assign_interval(const dpf_type & dpf, integral_type new_from, integral_type new_to)
|
||||
{
|
||||
static constexpr auto complement_of = std::bit_not{};
|
||||
|
|
@ -148,11 +158,11 @@ struct interval_memoizer_base
|
|||
std::size_t level_index; // indicates current level being built
|
||||
|
||||
explicit interval_memoizer_base(std::size_t output_len)
|
||||
: dpf_{std::nullopt},
|
||||
: output_length{output_len},
|
||||
level_index{0},
|
||||
dpf_{std::nullopt},
|
||||
from_{std::nullopt},
|
||||
to_{std::nullopt},
|
||||
output_length{output_len},
|
||||
level_index{0}
|
||||
to_{std::nullopt}
|
||||
{ }
|
||||
|
||||
private:
|
||||
|
|
@ -167,6 +177,7 @@ struct interval_memoizer_base
|
|||
/// `eval_interval(key, from, to)` allocates when you omit the memoizer.
|
||||
/// @tparam DpfKey DPF key type
|
||||
/// @tparam interior_node interior node
|
||||
/// \complexity O(L) nodes. The buffer length is about `interval_memoizer_slots(output_len)` (the last level plus the previous level's pivot). L is that output length.
|
||||
template <typename DpfKey,
|
||||
typename Allocator = aligned_allocator<
|
||||
typename interval_memoizer_key_t<DpfKey>::interior_node>>
|
||||
|
|
@ -248,6 +259,7 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
/// @brief Every level of the interval. `retains_all_levels` is true.
|
||||
/// @tparam DpfKey DPF key type
|
||||
/// @tparam interior_node interior node
|
||||
/// \complexity Allocates `level_endpoints[depth] + output_len` nodes: the prefix sum of `get_nodes_at_level` over every level, plus the output length.
|
||||
template <typename DpfKey,
|
||||
typename Allocator = aligned_allocator<
|
||||
typename interval_memoizer_key_t<DpfKey>::interior_node>>
|
||||
|
|
@ -326,7 +338,7 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
std::array<std::size_t, depth+1> level_endpoints{0};
|
||||
for (std::size_t level=depth; level > 0; --level)
|
||||
{
|
||||
len = std::min(len+2 >> 1, integral_type(1) << level-1);
|
||||
len = std::min((len + 2) >> 1, integral_type(1) << (level - 1));
|
||||
level_endpoints[level] = len;
|
||||
}
|
||||
|
||||
|
|
@ -468,10 +480,22 @@ HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
|||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||||
}
|
||||
|
||||
/// @brief Size for the tree walk of `[from, to]` after `offset_x`.
|
||||
/// @details Leaf packing depends on alignment (and wrap). Sizing on the
|
||||
/// logical endpoints alone can undersize once a wildcard `δ` is set.
|
||||
/// When the offset is not ready yet, falls back to logical endpoints
|
||||
/// (identity tree coordinates).
|
||||
template <typename DpfKey,
|
||||
typename InputT>
|
||||
inline auto make_basic_interval_memoizer(const DpfKey &, InputT from, InputT to)
|
||||
inline auto make_basic_interval_memoizer(const DpfKey & dpf, InputT from, InputT to)
|
||||
{
|
||||
using input_type = typename DpfKey::input_type;
|
||||
if (dpf.offset_x.is_ready())
|
||||
{
|
||||
return make_basic_interval_memoizer<DpfKey>(
|
||||
dpf.offset_x(static_cast<input_type>(from)),
|
||||
dpf.offset_x(static_cast<input_type>(to)));
|
||||
}
|
||||
return make_basic_interval_memoizer<DpfKey>(from, to);
|
||||
}
|
||||
|
||||
|
|
@ -513,8 +537,15 @@ HEDLEY_PRAGMA(GCC diagnostic pop)
|
|||
|
||||
template <typename DpfKey,
|
||||
typename InputT>
|
||||
inline auto make_full_tree_interval_memoizer(const DpfKey &, InputT from, InputT to)
|
||||
inline auto make_full_tree_interval_memoizer(const DpfKey & dpf, InputT from, InputT to)
|
||||
{
|
||||
using input_type = typename DpfKey::input_type;
|
||||
if (dpf.offset_x.is_ready())
|
||||
{
|
||||
return make_full_tree_interval_memoizer<DpfKey>(
|
||||
dpf.offset_x(static_cast<input_type>(from)),
|
||||
dpf.offset_x(static_cast<input_type>(to)));
|
||||
}
|
||||
return make_full_tree_interval_memoizer<DpfKey>(from, to);
|
||||
}
|
||||
|
||||
|
|
@ -583,8 +614,15 @@ inline auto make_basic_interval_memoizer(InputT from, InputT to)
|
|||
template <typename DpfKey, std::size_t I,
|
||||
typename InputT,
|
||||
std::enable_if_t<DpfKey::is_multilevel, bool> = true>
|
||||
inline auto make_basic_interval_memoizer(const DpfKey &, InputT from, InputT to)
|
||||
inline auto make_basic_interval_memoizer(const DpfKey & dpf, InputT from, InputT to)
|
||||
{
|
||||
using input_type = typename DpfKey::input_type;
|
||||
if (dpf.offset_x.is_ready())
|
||||
{
|
||||
return make_basic_interval_memoizer<DpfKey, I>(
|
||||
dpf.offset_x(static_cast<input_type>(from)),
|
||||
dpf.offset_x(static_cast<input_type>(to)));
|
||||
}
|
||||
return make_basic_interval_memoizer<DpfKey, I>(from, to);
|
||||
}
|
||||
|
||||
|
|
|
|||
94
include/dpf/it_dpf3.hpp
Normal file
94
include/dpf/it_dpf3.hpp
Normal file
|
|
@ -0,0 +1,94 @@
|
|||
/// @file dpf/it_dpf3.hpp
|
||||
/// @brief Information-theoretic 3-server distributed point function.
|
||||
/// @details Statistically private among three servers: any one key is
|
||||
/// independent of `(α, β)`, and the **sum** of all three evaluations
|
||||
/// is the point function (Boyle–Gilboa–Ishai–Kolobov, ePrint 2023/028).
|
||||
/// Distinct from the computational `(2,3)` Shamir key `make_dpf3`
|
||||
/// (ePrint 2024/1658), where any two shares reconstruct. Domain
|
||||
/// `uint8_t`, payload `uint64_t`. For this domain size the share is a
|
||||
/// full additive truth table (no GGM correction words); the paper's
|
||||
/// matching-vector packing targets asymptotically large `N`.
|
||||
/// @see dpf/dpf3.hpp, examples/applications/it_pir3.cpp, examples/applications/pir3.cpp
|
||||
/// @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_DPF_IT_DPF3_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_IT_DPF3_HPP__
|
||||
|
||||
#include <array>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <tuple>
|
||||
#include <utility>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/random.hpp"
|
||||
#include "dpf/secret_share.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
/// @brief One evaluator's key for the information-theoretic 3-server DPF.
|
||||
struct it_dpf3_key
|
||||
{
|
||||
static constexpr std::size_t domain_size = 256;
|
||||
using input_type = std::uint8_t;
|
||||
using output_type = std::uint64_t;
|
||||
|
||||
std::uint8_t party = 0;
|
||||
std::array<output_type, domain_size> share{};
|
||||
};
|
||||
|
||||
/// @brief Three additive keys for `f_{α,β}` on `{0..255} → uint64_t`.
|
||||
/// @details `eval_it_dpf3(k0,x) + eval_it_dpf3(k1,x) + eval_it_dpf3(k2,x)` equals
|
||||
/// `β` at `x = α` and `0` elsewhere (wrapping `uint64_t` arithmetic).
|
||||
/// \complexity O(N) samples for domain size `N = 256`. Key size is `N` words
|
||||
/// per party (truth-table share), not subpolynomial.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline auto make_it_dpf3(std::uint8_t alpha, std::uint64_t beta)
|
||||
{
|
||||
it_dpf3_key k0;
|
||||
it_dpf3_key k1;
|
||||
it_dpf3_key k2;
|
||||
k0.party = 0;
|
||||
k1.party = 1;
|
||||
k2.party = 2;
|
||||
for (std::size_t x = 0; x < it_dpf3_key::domain_size; ++x)
|
||||
{
|
||||
const std::uint64_t secret =
|
||||
(static_cast<std::uint8_t>(x) == alpha) ? beta : 0;
|
||||
auto [a, b, c] = additively_share3(secret);
|
||||
k0.share[x] = a.raw();
|
||||
k1.share[x] = b.raw();
|
||||
k2.share[x] = c.raw();
|
||||
}
|
||||
return std::make_tuple(std::move(k0), std::move(k1), std::move(k2));
|
||||
}
|
||||
|
||||
/// @brief One additive share of `f_{α,β}(x)`.
|
||||
/// \complexity O(1).
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
constexpr std::uint64_t eval_it_dpf3(const it_dpf3_key & key,
|
||||
std::uint8_t x) noexcept
|
||||
{
|
||||
return key.share[x];
|
||||
}
|
||||
|
||||
/// @brief `sum_x eval_it_dpf3(key, x) * weights[x]` over the whole domain.
|
||||
/// \complexity O(N) multiply-adds, `N = 256`.
|
||||
template <typename Weights>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
auto eval_it_dpf3_inner_product(const it_dpf3_key & key, Weights && weights)
|
||||
{
|
||||
std::uint64_t acc = 0;
|
||||
for (std::size_t x = 0; x < it_dpf3_key::domain_size; ++x)
|
||||
acc += key.share[x] * static_cast<std::uint64_t>(weights[x]);
|
||||
return acc;
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_IT_DPF3_HPP__
|
||||
|
|
@ -396,6 +396,8 @@ struct adl_serializer<dpf::beaver<true, NodeT, OutputT>, void>
|
|||
beaver.output_blind = dpf::json::codec::load<OutputT>(j.at("output_blind"));
|
||||
beaver.vector_blind = dpf::json::codec::load<typename beaver_type::LeafT>(j.at("vector_blind"));
|
||||
beaver.blinded_vector = dpf::json::codec::load<typename beaver_type::LeafT>(j.at("blinded_vector"));
|
||||
if (j.contains("assign_pad"))
|
||||
beaver.assign_pad = dpf::json::codec::load<typename beaver_type::LeafT>(j.at("assign_pad"));
|
||||
}
|
||||
|
||||
static void to_json(nlohmann::json & j, const dpf::beaver<true, NodeT, OutputT> & beaver) // NOLINT(runtime/references)
|
||||
|
|
@ -403,7 +405,8 @@ struct adl_serializer<dpf::beaver<true, NodeT, OutputT>, void>
|
|||
j = nlohmann::json{
|
||||
{"output_blind", dpf::json::codec::dump(beaver.output_blind)},
|
||||
{"vector_blind", dpf::json::codec::dump(beaver.vector_blind)},
|
||||
{"blinded_vector", dpf::json::codec::dump(beaver.blinded_vector)}
|
||||
{"blinded_vector", dpf::json::codec::dump(beaver.blinded_vector)},
|
||||
{"assign_pad", dpf::json::codec::dump(beaver.assign_pad)}
|
||||
};
|
||||
}
|
||||
};
|
||||
|
|
@ -506,7 +509,8 @@ struct adl_serializer<dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, OutputT, Ou
|
|||
entry["beaver"] = nlohmann::json{
|
||||
{"output_blind", dpf::json::codec::dump(beaver.output_blind)},
|
||||
{"vector_blind", dpf::json::codec::dump(beaver.vector_blind)},
|
||||
{"blinded_vector", dpf::json::codec::dump(beaver.blinded_vector)}
|
||||
{"blinded_vector", dpf::json::codec::dump(beaver.blinded_vector)},
|
||||
{"assign_pad", dpf::json::codec::dump(beaver.assign_pad)}
|
||||
};
|
||||
entry["output_share"] = dpf::json::codec::dump(wrapper.output_share());
|
||||
entry["state"] = wrapper.state();
|
||||
|
|
@ -535,6 +539,9 @@ struct adl_serializer<dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, OutputT, Ou
|
|||
stored.at("vector_blind"));
|
||||
beaver.blinded_vector = dpf::json::codec::load<decltype(beaver.blinded_vector)>(
|
||||
stored.at("blinded_vector"));
|
||||
if (stored.contains("assign_pad"))
|
||||
beaver.assign_pad = dpf::json::codec::load<decltype(beaver.assign_pad)>(
|
||||
stored.at("assign_pad"));
|
||||
return beaver;
|
||||
}
|
||||
else
|
||||
|
|
@ -604,6 +611,9 @@ struct adl_serializer<dpf::dpf_key<InteriorPRG, ExteriorPRG, InputT, OutputT, Ou
|
|||
stored.at("vector_blind"));
|
||||
beaver.blinded_vector = dpf::json::codec::load<decltype(beaver.blinded_vector)>(
|
||||
stored.at("blinded_vector"));
|
||||
if (stored.contains("assign_pad"))
|
||||
beaver.assign_pad = dpf::json::codec::load<decltype(beaver.assign_pad)>(
|
||||
stored.at("assign_pad"));
|
||||
wrapper out{std::move(leaf), std::move(beaver)};
|
||||
restore_leaf<I>(out, entry);
|
||||
return out;
|
||||
|
|
|
|||
|
|
@ -389,6 +389,10 @@ class basic_fixed_length_string : public dpf::modint<static_cast<std::size_t>(st
|
|||
/// uses `char` (i.e., bytes) as its **character type**, with its
|
||||
/// default `char_traits` and `allocator` types (see
|
||||
/// `dpf::basic_fixed_length_string` for more info on the template).
|
||||
/// @note A fixed-alphabet domain. The successor is `dpf::keyword2`, which
|
||||
/// stores a rank in a pattern language and is also not a leaf.
|
||||
/// @see dpf::keyword2
|
||||
/// @see dpf::alphabets
|
||||
template <std::size_t MaxLen,
|
||||
const char * Alphabet = alphabets::lowercase_alpha>
|
||||
using keyword = basic_fixed_length_string<MaxLen, char, Alphabet>;
|
||||
|
|
|
|||
|
|
@ -1686,7 +1686,10 @@ constexpr u256 unpack_rank(Integral v)
|
|||
} // namespace detail
|
||||
|
||||
/// @brief Ranked keyword whose language is the pattern `Pattern`.
|
||||
/// @note A domain. It stores a `modint` rank. It is not a leaf type.
|
||||
/// @tparam Pattern null-terminated pattern with static storage.
|
||||
/// @see dpf::keyword
|
||||
/// @see dpf::modint
|
||||
/// `pad('0')[0-9]{0,4}` is a decimal of at most 4 digits.
|
||||
/// `[0-9a-f]{8}` is exactly 8 significant hex digits.
|
||||
template <const char * Pattern>
|
||||
|
|
|
|||
62
include/dpf/launch.hpp
Normal file
62
include/dpf/launch.hpp
Normal file
|
|
@ -0,0 +1,62 @@
|
|||
/// @file dpf/launch.hpp
|
||||
/// @brief Start here: run a composed protocol as 2 or 3 parties.
|
||||
/// @details
|
||||
/// * In one process: `dpf::run_two_party` / `dpf::run_three_party`, with
|
||||
/// every choice in `app::run_config` (transport: async memory, unix
|
||||
/// sockets, TCP mux, parallel TCP, SCTP; lanes; framing; instances;
|
||||
/// window; timeouts). Party 2 is the dealer / RSS third party.
|
||||
/// * In separate processes: `app::parse_node_args(argc, argv)` then
|
||||
/// `app::run_node(args, plan, values)`, started as
|
||||
/// `--party=i --peers=host:port,host:port[,host:port]` on each machine.
|
||||
/// See `examples/protocol/party_node.cpp`.
|
||||
/// Both ends exchange a plan-shape hello, so a lane, framing, instance, or
|
||||
/// plan mismatch fails with a message that names the field.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_LAUNCH_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_LAUNCH_HPP__
|
||||
|
||||
#include <cstdint>
|
||||
#include <map>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/party_runner.hpp"
|
||||
#include "dpf/run_config.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
/// @brief Parties 0 and 1 on threads over `cfg.kind`.
|
||||
inline app::parties_result run_two_party(const protocol::plan & p0,
|
||||
const protocol::plan & p1, app::party_values & v0, app::party_values & v1,
|
||||
const std::map<std::uint32_t, protocol::kernel_fn> & kernels = {},
|
||||
const app::run_config & cfg = {})
|
||||
{
|
||||
std::vector<app::party_values> values(2);
|
||||
values[0] = std::move(v0);
|
||||
values[1] = std::move(v1);
|
||||
auto r = app::run_parties({p0, p1}, values, kernels, cfg);
|
||||
v0 = std::move(values[0]);
|
||||
v1 = std::move(values[1]);
|
||||
return r;
|
||||
}
|
||||
|
||||
/// @brief Parties 0, 1, and 2 (dealer / RSS third party) on threads.
|
||||
inline app::parties_result run_three_party(const protocol::plan & p0,
|
||||
const protocol::plan & p1, const protocol::plan & p2, app::party_values & v0,
|
||||
app::party_values & v1, app::party_values & v2,
|
||||
const std::map<std::uint32_t, protocol::kernel_fn> & kernels = {},
|
||||
const app::run_config & cfg = {})
|
||||
{
|
||||
std::vector<app::party_values> values(3);
|
||||
values[0] = std::move(v0);
|
||||
values[1] = std::move(v1);
|
||||
values[2] = std::move(v2);
|
||||
auto r = app::run_parties({p0, p1, p2}, values, kernels, cfg);
|
||||
v0 = std::move(values[0]);
|
||||
v1 = std::move(values[1]);
|
||||
v2 = std::move(values[2]);
|
||||
return r;
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif
|
||||
|
|
@ -23,6 +23,7 @@
|
|||
|
||||
#include "dpf/bit.hpp"
|
||||
#include "dpf/bitstring.hpp"
|
||||
#include "dpf/blob.hpp"
|
||||
#include "dpf/packed_lane_arithmetic.hpp"
|
||||
#include "dpf/wildcard.hpp"
|
||||
#include "dpf/xor_wrapper.hpp"
|
||||
|
|
@ -90,6 +91,83 @@ static constexpr auto subtract_leaf = leaf_arithmetic::subtract_t<OutputT, void>
|
|||
|
||||
static constexpr auto multiply_leaf = leaf_arithmetic::multiply_t<void, void>{};
|
||||
|
||||
/// @brief Scalar add in the leaf output group.
|
||||
/// @details `float`/`double` use XOR of the IEEE bit pattern (same as
|
||||
/// `add_t` on packed leaves), not IEEE floating-point addition.
|
||||
/// \complexity O(1) for a scalar. A packed node walks its bytes once (`add_t` / the lane loops): O(sizeof node) and the carry stays inside a `twobit` or `nyble` lane.
|
||||
template <typename T>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_PURE
|
||||
T leaf_group_add(const T & a, const T & b) noexcept
|
||||
{
|
||||
if constexpr (std::is_same_v<T, float> || std::is_same_v<T, double>)
|
||||
{
|
||||
using bits_t = std::conditional_t<sizeof(T) == 4, psnip_uint32_t, psnip_uint64_t>;
|
||||
bits_t aa{}, bb{};
|
||||
utils::raw_memcpy(&aa, &a, sizeof(T));
|
||||
utils::raw_memcpy(&bb, &b, sizeof(T));
|
||||
bits_t cc = static_cast<bits_t>(aa ^ bb);
|
||||
T out{};
|
||||
utils::raw_memcpy(&out, &cc, sizeof(T));
|
||||
return out;
|
||||
}
|
||||
else
|
||||
{
|
||||
return static_cast<T>(a + b);
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Scalar multiply in the leaf output group (`float`/`double`: AND of bits).
|
||||
template <typename T>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_PURE
|
||||
T leaf_group_mul(const T & a, const T & b) noexcept
|
||||
{
|
||||
if constexpr (std::is_same_v<T, float> || std::is_same_v<T, double>)
|
||||
{
|
||||
using bits_t = std::conditional_t<sizeof(T) == 4, psnip_uint32_t, psnip_uint64_t>;
|
||||
bits_t aa{}, bb{};
|
||||
utils::raw_memcpy(&aa, &a, sizeof(T));
|
||||
utils::raw_memcpy(&bb, &b, sizeof(T));
|
||||
bits_t cc = static_cast<bits_t>(aa & bb);
|
||||
T out{};
|
||||
utils::raw_memcpy(&out, &cc, sizeof(T));
|
||||
return out;
|
||||
}
|
||||
else
|
||||
{
|
||||
return static_cast<T>(a * b);
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Multiplicative identity for leaf scaling (all-ones for XOR/AND groups).
|
||||
template <typename T>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
HEDLEY_PURE
|
||||
T leaf_group_one() noexcept
|
||||
{
|
||||
if constexpr (std::is_same_v<T, float> || std::is_same_v<T, double>)
|
||||
{
|
||||
using bits_t = std::conditional_t<sizeof(T) == 4, psnip_uint32_t, psnip_uint64_t>;
|
||||
bits_t ones = static_cast<bits_t>(~bits_t{0});
|
||||
T out{};
|
||||
utils::raw_memcpy(&out, &ones, sizeof(T));
|
||||
return out;
|
||||
}
|
||||
else if constexpr (utils::is_xor_wrapper_v<T>)
|
||||
{
|
||||
using u = typename T::value_type;
|
||||
return T{static_cast<u>(~u{0})};
|
||||
}
|
||||
else
|
||||
{
|
||||
return T{1};
|
||||
}
|
||||
}
|
||||
|
||||
namespace leaf_arithmetic
|
||||
{
|
||||
|
||||
|
|
@ -260,10 +338,10 @@ struct add_array_t<OutputT, std::enable_if_t<dpf::utils::has_operators_plus_minu
|
|||
"arithmetic leaf array and output type must be the same size");
|
||||
std::array<T, N> c;
|
||||
output_type a_, b_;
|
||||
std::memcpy(&a_, std::data(a), sizeof(a_));
|
||||
std::memcpy(&b_, std::data(b), sizeof(b_));
|
||||
utils::raw_memcpy(&a_, std::data(a), sizeof(a_));
|
||||
utils::raw_memcpy(&b_, std::data(b), sizeof(b_));
|
||||
output_type c_ = a_ + b_;
|
||||
std::memcpy(std::data(c), &c_, sizeof(c_));
|
||||
utils::raw_memcpy(std::data(c), &c_, sizeof(c_));
|
||||
return c;
|
||||
}
|
||||
};
|
||||
|
|
@ -294,10 +372,10 @@ template <> struct add_t<simde_int128, simde__m128i> final
|
|||
{
|
||||
simde__m128i ret;
|
||||
simde_int128 lhs_, rhs_;
|
||||
std::memcpy(&lhs_, &lhs, sizeof(simde_int128));
|
||||
std::memcpy(&rhs_, &rhs, sizeof(simde_int128));
|
||||
utils::raw_memcpy(&lhs_, &lhs, sizeof(simde_int128));
|
||||
utils::raw_memcpy(&rhs_, &rhs, sizeof(simde_int128));
|
||||
simde_int128 sum = lhs_ + rhs_;
|
||||
std::memcpy(&ret, &sum, sizeof(simde__m128i));
|
||||
utils::raw_memcpy(&ret, &sum, sizeof(simde__m128i));
|
||||
return ret;
|
||||
}
|
||||
};
|
||||
|
|
@ -307,10 +385,10 @@ template <> struct add_t<simde_uint128, simde__m128i> final
|
|||
{
|
||||
simde__m128i ret;
|
||||
simde_uint128 lhs_, rhs_;
|
||||
std::memcpy(&lhs_, &lhs, sizeof(simde_uint128));
|
||||
std::memcpy(&rhs_, &rhs, sizeof(simde_uint128));
|
||||
utils::raw_memcpy(&lhs_, &lhs, sizeof(simde_uint128));
|
||||
utils::raw_memcpy(&rhs_, &rhs, sizeof(simde_uint128));
|
||||
simde_uint128 sum = lhs_ + rhs_;
|
||||
std::memcpy(&ret, &sum, sizeof(simde__m128i));
|
||||
utils::raw_memcpy(&ret, &sum, sizeof(simde__m128i));
|
||||
return ret;
|
||||
}
|
||||
};
|
||||
|
|
@ -335,10 +413,10 @@ template <> struct add_t<simde_int128, simde__m256i> final
|
|||
{
|
||||
simde__m256i ret;
|
||||
simde_int128 lhs_[2], rhs_[2];
|
||||
std::memcpy(&lhs_, &lhs, sizeof(simde_int128) * 2);
|
||||
std::memcpy(&rhs_, &rhs, sizeof(simde_int128) * 2);
|
||||
utils::raw_memcpy(&lhs_, &lhs, sizeof(simde_int128) * 2);
|
||||
utils::raw_memcpy(&rhs_, &rhs, sizeof(simde_int128) * 2);
|
||||
simde_int128 sum[2] = { lhs_[0] + rhs_[0], lhs_[1] + rhs_[1] };
|
||||
std::memcpy(&ret, &sum, sizeof(simde__m256i));
|
||||
utils::raw_memcpy(&ret, &sum, sizeof(simde__m256i));
|
||||
return ret;
|
||||
}
|
||||
};
|
||||
|
|
@ -348,21 +426,28 @@ template <> struct add_t<simde_uint128, simde__m256i> final
|
|||
{
|
||||
simde__m256i ret;
|
||||
simde_uint128 lhs_[2], rhs_[2];
|
||||
std::memcpy(&lhs_, &lhs, sizeof(simde_uint128) * 2);
|
||||
std::memcpy(&rhs_, &rhs, sizeof(simde_uint128) * 2);
|
||||
utils::raw_memcpy(&lhs_, &lhs, sizeof(simde_uint128) * 2);
|
||||
utils::raw_memcpy(&rhs_, &rhs, sizeof(simde_uint128) * 2);
|
||||
simde_uint128 sum[2] = { lhs_[0] + rhs_[0], lhs_[1] + rhs_[1] };
|
||||
std::memcpy(&ret, &sum, sizeof(simde__m256i));
|
||||
utils::raw_memcpy(&ret, &sum, sizeof(simde__m256i));
|
||||
return ret;
|
||||
}
|
||||
};
|
||||
|
||||
template <typename OutputT, typename NodeT, std::size_t N> struct add_t<OutputT, std::array<NodeT, N>> final : public detail::add_array_t<OutputT> {};
|
||||
template <std::size_t Nbits, typename WordT> struct add_t<dpf::bitstring<Nbits, WordT>, void> final : public detail::bitstring_xor_t<Nbits, WordT> {};
|
||||
template <std::size_t N> struct add_t<dpf::blob<N>, void> final : public detail::bitstring_xor_t<N * 8, unsigned char> {};
|
||||
template <std::size_t N, typename NodeT>
|
||||
struct add_t<dpf::blob<N>, NodeT> final : public detail::bitstring_xor_t<N * 8, unsigned char> {};
|
||||
/// @brief Bitwise XOR, not IEEE addition. Float addition does not form an
|
||||
/// exact secret-sharing group; XOR of the representation does.
|
||||
/// @tparam NodeT GGM node type
|
||||
template <typename NodeT> struct add_t<float, NodeT> final : public std::bit_xor<> {};
|
||||
template <typename NodeT> struct add_t<double, NodeT> final : public std::bit_xor<> {};
|
||||
/// @tparam NodeT GGM node type (not `void`; that selects the scalar wrapper)
|
||||
template <typename NodeT>
|
||||
struct add_t<float, NodeT, std::enable_if_t<!std::is_void_v<NodeT>>>
|
||||
final : public std::bit_xor<> {};
|
||||
template <typename NodeT>
|
||||
struct add_t<double, NodeT, std::enable_if_t<!std::is_void_v<NodeT>>>
|
||||
final : public std::bit_xor<> {};
|
||||
template <> struct add_t<dpf::bit, void> final : public std::bit_xor<> {};
|
||||
template <typename NodeT> struct add_t<dpf::bit, NodeT> final : public std::bit_xor<> {};
|
||||
template <typename T> struct add_t<xor_wrapper<T>, void> final : public std::bit_xor<> {};
|
||||
|
|
@ -546,10 +631,10 @@ struct sub_array_t<OutputT, std::enable_if_t<dpf::utils::has_operators_plus_minu
|
|||
"arithmetic leaf array and output type must be the same size");
|
||||
std::array<T, N> c;
|
||||
output_type a_, b_;
|
||||
std::memcpy(&a_, std::data(a), sizeof(a_));
|
||||
std::memcpy(&b_, std::data(b), sizeof(b_));
|
||||
utils::raw_memcpy(&a_, std::data(a), sizeof(a_));
|
||||
utils::raw_memcpy(&b_, std::data(b), sizeof(b_));
|
||||
output_type c_ = a_ - b_;
|
||||
std::memcpy(std::data(c), &c_, sizeof(c_));
|
||||
utils::raw_memcpy(std::data(c), &c_, sizeof(c_));
|
||||
return c;
|
||||
}
|
||||
};
|
||||
|
|
@ -580,10 +665,10 @@ template <> struct subtract_t<simde_int128, simde__m128i> final
|
|||
{
|
||||
simde__m128i ret;
|
||||
simde_int128 lhs_, rhs_;
|
||||
std::memcpy(&lhs_, &lhs, sizeof(simde_int128));
|
||||
std::memcpy(&rhs_, &rhs, sizeof(simde_int128));
|
||||
utils::raw_memcpy(&lhs_, &lhs, sizeof(simde_int128));
|
||||
utils::raw_memcpy(&rhs_, &rhs, sizeof(simde_int128));
|
||||
simde_int128 sum = lhs_ - rhs_;
|
||||
std::memcpy(&ret, &sum, sizeof(simde__m128i));
|
||||
utils::raw_memcpy(&ret, &sum, sizeof(simde__m128i));
|
||||
return ret;
|
||||
}
|
||||
};
|
||||
|
|
@ -593,10 +678,10 @@ template <> struct subtract_t<simde_uint128, simde__m128i> final
|
|||
{
|
||||
simde__m128i ret;
|
||||
simde_uint128 lhs_, rhs_;
|
||||
std::memcpy(&lhs_, &lhs, sizeof(simde_uint128));
|
||||
std::memcpy(&rhs_, &rhs, sizeof(simde_uint128));
|
||||
utils::raw_memcpy(&lhs_, &lhs, sizeof(simde_uint128));
|
||||
utils::raw_memcpy(&rhs_, &rhs, sizeof(simde_uint128));
|
||||
simde_uint128 sum = lhs_ - rhs_;
|
||||
std::memcpy(&ret, &sum, sizeof(simde__m128i));
|
||||
utils::raw_memcpy(&ret, &sum, sizeof(simde__m128i));
|
||||
return ret;
|
||||
}
|
||||
};
|
||||
|
|
@ -622,10 +707,10 @@ template <> struct subtract_t<simde_int128, simde__m256i> final
|
|||
{
|
||||
simde__m256i ret;
|
||||
simde_int128 lhs_[2], rhs_[2];
|
||||
std::memcpy(&lhs_, &lhs, sizeof(simde_int128) * 2);
|
||||
std::memcpy(&rhs_, &rhs, sizeof(simde_int128) * 2);
|
||||
utils::raw_memcpy(&lhs_, &lhs, sizeof(simde_int128) * 2);
|
||||
utils::raw_memcpy(&rhs_, &rhs, sizeof(simde_int128) * 2);
|
||||
simde_int128 sum[2] = { lhs_[0] - rhs_[0], lhs_[1] - rhs_[1] };
|
||||
std::memcpy(&ret, &sum, sizeof(simde__m256i));
|
||||
utils::raw_memcpy(&ret, &sum, sizeof(simde__m256i));
|
||||
return ret;
|
||||
}
|
||||
};
|
||||
|
|
@ -635,20 +720,27 @@ template <> struct subtract_t<simde_uint128, simde__m256i> final
|
|||
{
|
||||
simde__m256i ret;
|
||||
simde_uint128 lhs_[2], rhs_[2];
|
||||
std::memcpy(&lhs_, &lhs, sizeof(simde_uint128) * 2);
|
||||
std::memcpy(&rhs_, &rhs, sizeof(simde_uint128) * 2);
|
||||
utils::raw_memcpy(&lhs_, &lhs, sizeof(simde_uint128) * 2);
|
||||
utils::raw_memcpy(&rhs_, &rhs, sizeof(simde_uint128) * 2);
|
||||
simde_uint128 sum[2] = { lhs_[0] - rhs_[0], lhs_[1] - rhs_[1] };
|
||||
std::memcpy(&ret, &sum, sizeof(simde__m256i));
|
||||
utils::raw_memcpy(&ret, &sum, sizeof(simde__m256i));
|
||||
return ret;
|
||||
}
|
||||
};
|
||||
|
||||
template <typename OutputT, typename NodeT, std::size_t N> struct subtract_t<OutputT, std::array<NodeT, N>> final : public detail::sub_array_t<OutputT> {};
|
||||
template <std::size_t Nbits, typename WordT> struct subtract_t<dpf::bitstring<Nbits, WordT>, void> final : public detail::bitstring_xor_t<Nbits, WordT> {};
|
||||
template <std::size_t N> struct subtract_t<dpf::blob<N>, void> final : public detail::bitstring_xor_t<N * 8, unsigned char> {};
|
||||
template <std::size_t N, typename NodeT>
|
||||
struct subtract_t<dpf::blob<N>, NodeT> final : public detail::bitstring_xor_t<N * 8, unsigned char> {};
|
||||
/// @brief Bitwise XOR, not IEEE subtraction.
|
||||
/// @tparam NodeT GGM node type
|
||||
template <typename NodeT> struct subtract_t<float, NodeT> final : public std::bit_xor<> {};
|
||||
template <typename NodeT> struct subtract_t<double, NodeT> final : public std::bit_xor<> {};
|
||||
template <typename NodeT>
|
||||
struct subtract_t<float, NodeT, std::enable_if_t<!std::is_void_v<NodeT>>>
|
||||
final : public std::bit_xor<> {};
|
||||
template <typename NodeT>
|
||||
struct subtract_t<double, NodeT, std::enable_if_t<!std::is_void_v<NodeT>>>
|
||||
final : public std::bit_xor<> {};
|
||||
template <typename NodeT> struct subtract_t<dpf::bit, NodeT> final : public std::bit_xor<> {};
|
||||
template <> struct subtract_t<dpf::bit, void> final : public std::bit_xor<> {};
|
||||
template <typename T> struct subtract_t<xor_wrapper<T>, void> final : public std::bit_xor<> {};
|
||||
|
|
@ -847,9 +939,9 @@ template <> struct multiply_t<simde_int128, simde__m128i> final
|
|||
{
|
||||
simde_int128 a_;
|
||||
simde__m128i c;
|
||||
std::memcpy(&a_, &a, sizeof(simde_int128));
|
||||
utils::raw_memcpy(&a_, &a, sizeof(simde_int128));
|
||||
simde_int128 c_ = a_ * b;
|
||||
std::memcpy(&c, &c_, sizeof(simde__m128i));
|
||||
utils::raw_memcpy(&c, &c_, sizeof(simde__m128i));
|
||||
return c;
|
||||
}
|
||||
};
|
||||
|
|
@ -860,9 +952,9 @@ template <> struct multiply_t<simde_uint128, simde__m128i> final
|
|||
{
|
||||
simde_uint128 a_;
|
||||
simde__m128i c;
|
||||
std::memcpy(&a_, &a, sizeof(simde_uint128));
|
||||
utils::raw_memcpy(&a_, &a, sizeof(simde_uint128));
|
||||
simde_uint128 c_ = a_ * b;
|
||||
std::memcpy(&c, &c_, sizeof(simde__m128i));
|
||||
utils::raw_memcpy(&c, &c_, sizeof(simde__m128i));
|
||||
return c;
|
||||
}
|
||||
};
|
||||
|
|
@ -908,8 +1000,8 @@ struct multiply_t<xor_wrapper<T>, simde__m128i> final
|
|||
} else {
|
||||
alignas(simde__m128i) unsigned char buf[sizeof(simde__m128i)]{};
|
||||
for (std::size_t off = 0; off + sizeof(val_t) <= sizeof(buf); off += sizeof(val_t))
|
||||
std::memcpy(buf + off, &v, sizeof(val_t));
|
||||
std::memcpy(&bb, buf, sizeof(bb));
|
||||
utils::raw_memcpy(buf + off, &v, sizeof(val_t));
|
||||
utils::raw_memcpy(&bb, buf, sizeof(bb));
|
||||
}
|
||||
return simde_mm_and_si128(a, bb);
|
||||
}
|
||||
|
|
@ -933,8 +1025,8 @@ struct multiply_t<xor_wrapper<T>, simde__m256i> final
|
|||
} else {
|
||||
alignas(simde__m256i) unsigned char buf[sizeof(simde__m256i)]{};
|
||||
for (std::size_t off = 0; off + sizeof(val_t) <= sizeof(buf); off += sizeof(val_t))
|
||||
std::memcpy(buf + off, &v, sizeof(val_t));
|
||||
std::memcpy(&bb, buf, sizeof(bb));
|
||||
utils::raw_memcpy(buf + off, &v, sizeof(val_t));
|
||||
utils::raw_memcpy(&bb, buf, sizeof(bb));
|
||||
}
|
||||
return simde_mm256_and_si256(a, bb);
|
||||
}
|
||||
|
|
@ -986,7 +1078,7 @@ struct multiply_t<float, simde__m128i> final
|
|||
{
|
||||
static_assert(sizeof(float) == 4, "float must be 32 bits");
|
||||
psnip_uint32_t bits = 0;
|
||||
std::memcpy(&bits, &b, sizeof(bits));
|
||||
utils::raw_memcpy(&bits, &b, sizeof(bits));
|
||||
return simde_mm_and_si128(a, simde_mm_set1_epi32(static_cast<int>(bits)));
|
||||
}
|
||||
};
|
||||
|
|
@ -998,7 +1090,7 @@ struct multiply_t<float, simde__m256i> final
|
|||
{
|
||||
static_assert(sizeof(float) == 4, "float must be 32 bits");
|
||||
psnip_uint32_t bits = 0;
|
||||
std::memcpy(&bits, &b, sizeof(bits));
|
||||
utils::raw_memcpy(&bits, &b, sizeof(bits));
|
||||
return simde_mm256_and_si256(a, simde_mm256_set1_epi32(static_cast<int>(bits)));
|
||||
}
|
||||
};
|
||||
|
|
@ -1010,7 +1102,7 @@ struct multiply_t<double, simde__m128i> final
|
|||
{
|
||||
static_assert(sizeof(double) == 8, "double must be 64 bits");
|
||||
psnip_uint64_t bits = 0;
|
||||
std::memcpy(&bits, &b, sizeof(bits));
|
||||
utils::raw_memcpy(&bits, &b, sizeof(bits));
|
||||
return simde_mm_and_si128(a, simde_mm_set1_epi64x(static_cast<long long>(bits)));
|
||||
}
|
||||
};
|
||||
|
|
@ -1022,7 +1114,7 @@ struct multiply_t<double, simde__m256i> final
|
|||
{
|
||||
static_assert(sizeof(double) == 8, "double must be 64 bits");
|
||||
psnip_uint64_t bits = 0;
|
||||
std::memcpy(&bits, &b, sizeof(bits));
|
||||
utils::raw_memcpy(&bits, &b, sizeof(bits));
|
||||
return simde_mm256_and_si256(a, simde_mm256_set1_epi64x(static_cast<long long>(bits)));
|
||||
}
|
||||
};
|
||||
|
|
|
|||
232
include/dpf/leaf_later.hpp
Normal file
232
include/dpf/leaf_later.hpp
Normal file
|
|
@ -0,0 +1,232 @@
|
|||
/// @file dpf/leaf_later.hpp
|
||||
/// @brief Defer the leaf correction until a public group element is known.
|
||||
/// @details `eval_full` / `eval_full_add_into` with `leaf_later` skip the leaf
|
||||
/// correction word, write the uncorrected group share, and fill a
|
||||
/// parallel control-bit buffer the caller owns. After a rotate of
|
||||
/// value and control together, `apply_leaf_correction` does
|
||||
/// `buf[i] += F * control[i]` so Duoram's `v[i−s] + F·t[i−s]` never
|
||||
/// reapplies a walk.
|
||||
/// @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_DPF_LEAF_LATER_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_LEAF_LATER_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <limits>
|
||||
#include <stdexcept>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/dpf_key.hpp"
|
||||
#include "dpf/eval_common.hpp"
|
||||
#include "dpf/eval_interval.hpp"
|
||||
#include "dpf/interval_memoizer.hpp"
|
||||
#include "dpf/leaf_arithmetic.hpp"
|
||||
#include "dpf/leaf_node.hpp"
|
||||
#include "dpf/twiddle.hpp"
|
||||
#include "dpf/utils.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
/// @brief Tag: expand without applying the leaf correction word.
|
||||
struct leaf_later
|
||||
{
|
||||
static constexpr bool is_leaf_later_tag = true;
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
struct is_leaf_later : std::false_type {};
|
||||
template <>
|
||||
struct is_leaf_later<leaf_later> : std::true_type {};
|
||||
template <typename T>
|
||||
inline constexpr bool is_leaf_later_v = is_leaf_later<std::decay_t<T>>::value;
|
||||
|
||||
// Forward declaration: defined in eval_walk.hpp (same offset convention).
|
||||
struct rotate;
|
||||
|
||||
namespace detail_leaf_later
|
||||
{
|
||||
|
||||
template <typename KeyT>
|
||||
HEDLEY_NO_THROW
|
||||
constexpr std::size_t domain_size() noexcept
|
||||
{
|
||||
constexpr std::size_t bits =
|
||||
utils::bitlength_of_v<typename KeyT::input_type>;
|
||||
static_assert(bits < 8 * sizeof(std::size_t),
|
||||
"leaf_later: input domain does not fit a std::size_t index");
|
||||
return std::size_t{1} << bits;
|
||||
}
|
||||
|
||||
/// @brief Uncorrected exterior leaf: `0 − mask` (same group as traverse_exterior).
|
||||
template <std::size_t I = 0, typename DpfKey>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
auto traverse_exterior_uncorrected(const typename DpfKey::interior_node & node)
|
||||
{
|
||||
using output_type = typename DpfKey::concrete_output_type<I>;
|
||||
using exterior_prg = typename DpfKey::exterior_prg;
|
||||
using outputs_tuple = typename DpfKey::concrete_outputs_tuple;
|
||||
using leaf_type = dpf::leaf_node_t<typename DpfKey::exterior_node, output_type>;
|
||||
leaf_type zero{};
|
||||
return dpf::subtract_leaf<output_type>(
|
||||
zero,
|
||||
make_leaf_mask_inner<exterior_prg, I, outputs_tuple>(
|
||||
unset_lo_2bits(node)));
|
||||
}
|
||||
|
||||
template <typename DpfKey, typename = void>
|
||||
struct has_opl_of : std::false_type {};
|
||||
template <typename DpfKey>
|
||||
struct has_opl_of<DpfKey,
|
||||
std::void_t<decltype(DpfKey::template outputs_per_leaf_of<0>)>>
|
||||
: std::true_type {};
|
||||
|
||||
template <typename DpfKey>
|
||||
HEDLEY_NO_THROW
|
||||
constexpr std::size_t opl_of() noexcept
|
||||
{
|
||||
if constexpr (has_opl_of<DpfKey>::value)
|
||||
return DpfKey::template outputs_per_leaf_of<0>;
|
||||
else
|
||||
return DpfKey::outputs_per_leaf;
|
||||
}
|
||||
|
||||
} // namespace detail_leaf_later
|
||||
|
||||
/// @brief Full-domain expansion without the leaf correction; also fills `control`.
|
||||
/// @details `buf[i]` receives the uncorrected share at domain point `i`, and
|
||||
/// `control[i]` is the leaf node's control bit. Both must be sized to
|
||||
/// `2^n`.
|
||||
/// \complexity One full-domain expansion, `Θ(2^n)`.
|
||||
template <std::size_t I = 0, typename Buffer, typename Control, typename DpfKey,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
void eval_full(Buffer & buf, Control & control, const DpfKey & dpf,
|
||||
leaf_later) // NOLINT(runtime/references)
|
||||
{
|
||||
using input_type = typename DpfKey::input_type;
|
||||
using output_type = typename DpfKey::concrete_output_type<I>;
|
||||
using exterior_node = typename DpfKey::exterior_node;
|
||||
constexpr std::size_t opl = DpfKey::outputs_per_leaf;
|
||||
const std::size_t n = detail_leaf_later::domain_size<DpfKey>();
|
||||
if (buf.size() < n || control.size() < n)
|
||||
throw std::invalid_argument("eval_full(leaf_later): buffer too small");
|
||||
|
||||
auto memo = make_basic_full_memoizer(dpf);
|
||||
const input_type from = std::numeric_limits<input_type>::min();
|
||||
const input_type to = std::numeric_limits<input_type>::max();
|
||||
const auto from_node = utils::get_from_node<DpfKey>(from);
|
||||
const auto to_node = utils::get_to_node<DpfKey>(to);
|
||||
internal::eval_interval_interior(dpf, from_node, to_node, memo);
|
||||
|
||||
const std::size_t nodes_in_interval =
|
||||
static_cast<std::size_t>(to_node - from_node);
|
||||
auto * nodes = memo[DpfKey::depth];
|
||||
|
||||
for (std::size_t j = 0; j < nodes_in_interval; ++j)
|
||||
{
|
||||
const auto & node = nodes[j];
|
||||
const bool tbit = static_cast<bool>(get_lo_bit(node));
|
||||
auto leaf = detail_leaf_later::traverse_exterior_uncorrected<I, DpfKey>(
|
||||
node);
|
||||
for (std::size_t o = 0; o < opl; ++o)
|
||||
{
|
||||
const std::size_t idx = j * opl + o;
|
||||
if (idx >= n)
|
||||
break;
|
||||
output_type y = extract_leaf<exterior_node, output_type>(leaf, o);
|
||||
buf[idx] = y;
|
||||
control[idx] = tbit;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Add an uncorrected full-domain expansion into `buf` and fill `control`.
|
||||
template <std::size_t I = 0, typename Buffer, typename Control, typename DpfKey,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
void eval_full_add_into(Buffer & buf, Control & control, const DpfKey & dpf,
|
||||
leaf_later) // NOLINT(runtime/references)
|
||||
{
|
||||
using output_type = typename DpfKey::concrete_output_type<I>;
|
||||
const std::size_t n = detail_leaf_later::domain_size<DpfKey>();
|
||||
std::vector<output_type> tmp(n);
|
||||
std::vector<std::uint8_t> ctl(n);
|
||||
eval_full<I>(tmp, ctl, dpf, leaf_later{});
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
buf[i] = buf[i] + tmp[i];
|
||||
control[i] = static_cast<bool>(ctl[i]);
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Uncorrected expansion written at `(i + rot.shift) mod 2^n`, with
|
||||
/// control bits rotated the same way.
|
||||
template <std::size_t I = 0, typename Buffer, typename Control, typename DpfKey,
|
||||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||||
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
||||
void eval_full_add_into(Buffer & buf, Control & control, const DpfKey & dpf,
|
||||
leaf_later, rotate rot) // NOLINT(runtime/references)
|
||||
{
|
||||
using output_type = typename DpfKey::concrete_output_type<I>;
|
||||
const std::size_t n = detail_leaf_later::domain_size<DpfKey>();
|
||||
const std::size_t s = rot.shift % n;
|
||||
std::vector<output_type> tmp(n);
|
||||
std::vector<std::uint8_t> ctl(n);
|
||||
eval_full<I>(tmp, ctl, dpf, leaf_later{});
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
const std::size_t j = (i + s) % n;
|
||||
buf[j] = buf[j] + tmp[i];
|
||||
control[j] = static_cast<bool>(ctl[i]);
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Rotate value and control by the same offset (`new[i] = old[(i-s) mod n]`).
|
||||
/// @details Domain point `k` lands at `(k + s) mod n`, matching `dpf::rotate{s}`.
|
||||
template <typename Buffer, typename Control>
|
||||
void cyclic_shift_pair(Buffer & buf, Control & control, std::size_t s) // NOLINT(runtime/references)
|
||||
{
|
||||
const std::size_t n = buf.size();
|
||||
if (n == 0 || control.size() != n)
|
||||
throw std::invalid_argument("cyclic_shift_pair: size mismatch");
|
||||
const std::size_t sh = s % n;
|
||||
if (sh == 0)
|
||||
return;
|
||||
Buffer out_buf(buf);
|
||||
Control out_ctl(control);
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
out_buf[i] = buf[(i + n - sh) % n];
|
||||
out_ctl[i] = control[(i + n - sh) % n];
|
||||
}
|
||||
buf = std::move(out_buf);
|
||||
control = std::move(out_ctl);
|
||||
}
|
||||
|
||||
/// @brief `buf[i] += F * control[i]` in the leaf's group.
|
||||
/// @details XOR leaves treat `+` as XOR. `F` is a public group element.
|
||||
template <typename Buffer, typename Control, typename F>
|
||||
void apply_leaf_correction(Buffer & buf, const Control & control, const F & f) // NOLINT(runtime/references)
|
||||
{
|
||||
const std::size_t n = buf.size();
|
||||
if (control.size() < n)
|
||||
throw std::invalid_argument("apply_leaf_correction: control too small");
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
if (static_cast<bool>(control[i]))
|
||||
buf[i] = buf[i] + f;
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_LEAF_LATER_HPP__
|
||||
|
|
@ -77,6 +77,27 @@ template <typename OutputT,
|
|||
static constexpr std::size_t block_length_of_leaf_v
|
||||
= block_length_of_leaf<OutputT, NodeT>::value;
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// dpf::blob<N> stretched byte leaves.
|
||||
// A `dpf::blob<N>` is an XOR share of `N` bytes and is *never* packed several
|
||||
// per node: Boyle packing keeps `lg(outputs_per_leaf) = 0` so the tree depth
|
||||
// follows the input bitlength (Express `genDPF`). Instead it stretches the
|
||||
// final seed over `ceil(8N / lambda)` exterior blocks in counter mode. These
|
||||
// specializations override the generic `is_packable` heuristic, which would
|
||||
// otherwise pack small blobs (e.g. `blob<1>`, whose 8 bits divide `lambda`).
|
||||
// See dpf/blob.hpp.
|
||||
template <std::size_t N,
|
||||
typename NodeT>
|
||||
struct outputs_per_leaf<dpf::blob<N>, NodeT>
|
||||
: public std::integral_constant<std::size_t, 1> { };
|
||||
|
||||
template <std::size_t N,
|
||||
typename NodeT>
|
||||
struct block_length_of_leaf<dpf::blob<N>, NodeT>
|
||||
: public std::integral_constant<std::size_t,
|
||||
utils::quotient_ceiling(N * 8,
|
||||
utils::bitlength_of_output_v<NodeT, NodeT>)> { };
|
||||
|
||||
template <typename OutputT,
|
||||
typename NodeT,
|
||||
typename InputT>
|
||||
|
|
@ -205,6 +226,10 @@ struct beaver<true, NodeT, OutputT> final
|
|||
OutputT output_blind;
|
||||
LeafT vector_blind;
|
||||
LeafT blinded_vector;
|
||||
/// @brief Keygen pad `vector_blind_i · output_blind_{1-i}` so a later
|
||||
/// `begin_update` can open a fresh naked delta without replaying
|
||||
/// the zero-payload correction word.
|
||||
LeafT assign_pad{};
|
||||
};
|
||||
|
||||
template <typename NodeT,
|
||||
|
|
@ -287,6 +312,33 @@ constexpr auto * leaf_blocks(LeafT & leaf) noexcept
|
|||
return leaf.data();
|
||||
}
|
||||
|
||||
/// @brief After a PRG fills `leaf`, map curve-point types through `from_seed`.
|
||||
template <typename Concrete, typename = void>
|
||||
struct is_curve_leaf_encode : std::false_type {};
|
||||
|
||||
template <typename Concrete>
|
||||
struct is_curve_leaf_encode<Concrete, std::void_t<
|
||||
std::bool_constant<Concrete::dpf_curve_point>,
|
||||
decltype(Concrete::from_seed(static_cast<const void *>(nullptr),
|
||||
std::size_t{0})),
|
||||
std::integral_constant<std::size_t, Concrete::encoded_size>,
|
||||
decltype(std::declval<const Concrete &>().bytes())>>
|
||||
: std::bool_constant<Concrete::dpf_curve_point> {};
|
||||
|
||||
template <typename Concrete, typename Leaf>
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void encode_curve_leaf_mask(Leaf & leaf) noexcept
|
||||
{
|
||||
if constexpr (is_curve_leaf_encode<Concrete>::value)
|
||||
{
|
||||
const Concrete pt = Concrete::from_seed(
|
||||
std::addressof(leaf), sizeof(leaf));
|
||||
leaf = Leaf{};
|
||||
std::memcpy(std::addressof(leaf), pt.bytes(),
|
||||
Concrete::encoded_size);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename ExteriorPRG,
|
||||
std::size_t I,
|
||||
typename OutputsTuple,
|
||||
|
|
@ -295,6 +347,7 @@ auto make_leaf_mask_inner(const InteriorBlock & seed, std::size_t pos_base = 0)
|
|||
{
|
||||
using node_type = typename ExteriorPRG::block_type;
|
||||
using output_type = std::tuple_element_t<I, OutputsTuple>;
|
||||
using concrete = concrete_type_t<output_type>;
|
||||
HEDLEY_PRAGMA(GCC diagnostic push)
|
||||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
||||
using leaf_type = dpf::leaf_node_t<node_type, output_type>;
|
||||
|
|
@ -305,6 +358,7 @@ HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
|||
auto seed_ = utils::to_exterior_node<node_type>(seed);
|
||||
ExteriorPRG::eval(seed_, leaf_blocks<node_type>(output), count,
|
||||
static_cast<psnip_uint32_t>(pos));
|
||||
encode_curve_leaf_mask<concrete>(output);
|
||||
|
||||
return output;
|
||||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||||
|
|
@ -370,6 +424,23 @@ auto make_leaves_impl(InputT x, const ExteriorBlock & seed0, const ExteriorBlock
|
|||
make_leaf<ExteriorPRG, Is>(x, seed0, seed1, sign, pos_base, ys...)...);
|
||||
}
|
||||
|
||||
namespace beavers
|
||||
{
|
||||
|
||||
/// @brief Fill wildcard leaf scale blinds from `sample_scale`.
|
||||
/// @tparam Concrete lane type
|
||||
/// @tparam LeafT packed leaf type
|
||||
/// @tparam Sample ring sampler
|
||||
template <typename Concrete, typename LeafT, typename Sample>
|
||||
void fill_wildcard_scale_blinds(Concrete & out0, Concrete & out1, LeafT & vec0,
|
||||
LeafT & vec1, std::size_t nlanes, Sample && rng);
|
||||
|
||||
template <typename Concrete, typename LeafT>
|
||||
void fill_wildcard_scale_blinds(Concrete & out0, Concrete & out1, LeafT & vec0,
|
||||
LeafT & vec1, std::size_t nlanes);
|
||||
|
||||
} // namespace beavers
|
||||
|
||||
template <typename ExteriorPRG,
|
||||
typename InputT,
|
||||
typename ExteriorBlock,
|
||||
|
|
@ -435,38 +506,43 @@ HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
|||
// secret share the value
|
||||
dpf::uniform_fill(leaf0);
|
||||
leaf1 = dpf::subtract_leaf<concrete_type>(leaf, leaf0);
|
||||
// also initialize the beavers
|
||||
if constexpr(!dpf::utils::has_characteristic_two_v<concrete_type>
|
||||
|| dpf::outputs_per_leaf_v<concrete_type, node_type> > 1)
|
||||
// Always plant a scale Beaver, including full-width XOR
|
||||
// (char-2, one lane). Skipping it left blinds at zero so
|
||||
// online assign exchanged beta in the clear.
|
||||
dpf::leaf_node_t<node_type, concrete_type> vector;
|
||||
// XOR/AND leaf groups use the all-ones word as unit, not ±1.
|
||||
// Check the OUTPUT type: input may be modint while the
|
||||
// leaf is xor_wrapper (wildcard XOR payload). IEEE
|
||||
// float/double leaves are the same bitwise group.
|
||||
if constexpr(utils::is_xor_wrapper_v<std::decay_t<decltype(x)>> == true
|
||||
|| utils::is_xor_wrapper_v<concrete_type> == true
|
||||
|| std::is_same_v<concrete_type, float>
|
||||
|| std::is_same_v<concrete_type, double>)
|
||||
{
|
||||
dpf::leaf_node_t<node_type, concrete_type> vector;
|
||||
// XOR-group multiply is AND, whose unit is ~0, not ±1.
|
||||
// Check the OUTPUT type: input may be modint while the
|
||||
// leaf is xor_wrapper (wildcard XOR payload).
|
||||
if constexpr(utils::is_xor_wrapper_v<std::decay_t<decltype(x)>> == true
|
||||
|| utils::is_xor_wrapper_v<concrete_type> == true)
|
||||
{
|
||||
vector = make_naked_leaf<node_type>(x, concrete_type(~0));
|
||||
}
|
||||
else
|
||||
{
|
||||
vector = make_naked_leaf<node_type>(x, concrete_type(2*sign-1));
|
||||
}
|
||||
|
||||
uniform_fill(beaver0.output_blind);
|
||||
uniform_fill(beaver0.vector_blind);
|
||||
|
||||
uniform_fill(beaver1.output_blind);
|
||||
uniform_fill(beaver1.vector_blind);
|
||||
|
||||
beaver0.blinded_vector = dpf::add_leaf<concrete_type>(vector, beaver1.vector_blind);
|
||||
beaver1.blinded_vector = dpf::add_leaf<concrete_type>(vector, beaver0.vector_blind);
|
||||
|
||||
leaf0 = dpf::add_leaf<concrete_type>(leaf0,
|
||||
dpf::multiply_leaf(beaver0.vector_blind, beaver1.output_blind));
|
||||
leaf1 = dpf::add_leaf<concrete_type>(leaf1,
|
||||
dpf::multiply_leaf(beaver1.vector_blind, beaver0.output_blind));
|
||||
vector = make_naked_leaf<node_type>(x,
|
||||
dpf::leaf_group_one<concrete_type>());
|
||||
}
|
||||
else
|
||||
{
|
||||
vector = make_naked_leaf<node_type>(x, concrete_type(2*sign-1));
|
||||
}
|
||||
|
||||
constexpr std::size_t nlanes =
|
||||
dpf::outputs_per_leaf_v<concrete_type, node_type>;
|
||||
dpf::beavers::fill_wildcard_scale_blinds(
|
||||
beaver0.output_blind, beaver1.output_blind,
|
||||
beaver0.vector_blind, beaver1.vector_blind,
|
||||
nlanes > 0 ? nlanes : std::size_t{1});
|
||||
|
||||
beaver0.blinded_vector = dpf::add_leaf<concrete_type>(vector, beaver1.vector_blind);
|
||||
beaver1.blinded_vector = dpf::add_leaf<concrete_type>(vector, beaver0.vector_blind);
|
||||
|
||||
beaver0.assign_pad = dpf::multiply_leaf(
|
||||
beaver0.vector_blind, beaver1.output_blind);
|
||||
beaver1.assign_pad = dpf::multiply_leaf(
|
||||
beaver1.vector_blind, beaver0.output_blind);
|
||||
leaf0 = dpf::add_leaf<concrete_type>(leaf0, beaver0.assign_pad);
|
||||
leaf1 = dpf::add_leaf<concrete_type>(leaf1, beaver1.assign_pad);
|
||||
}
|
||||
else
|
||||
{
|
||||
|
|
|
|||
|
|
@ -8,13 +8,20 @@
|
|||
#ifndef LIBDPF_INCLUDE_DPF_LEAF_WRAPPER_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_LEAF_WRAPPER_HPP__
|
||||
|
||||
#include <stdexcept>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/leaf_arithmetic.hpp"
|
||||
#include "dpf/secret_share.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
|
||||
/// @brief Concrete leaf. `get()` is ready immediately.
|
||||
/// @tparam OutputT output type
|
||||
/// @tparam NodeT exterior node type
|
||||
/// @see dpf::wildcard_value
|
||||
template <typename OutputT,
|
||||
typename NodeT>
|
||||
struct leaf_wrapper
|
||||
|
|
@ -31,10 +38,18 @@ struct leaf_wrapper
|
|||
HEDLEY_NO_THROW
|
||||
constexpr const leaf_type & get() const noexcept { return leaf_; }
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
constexpr leaf_type & get() noexcept { return leaf_; }
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
constexpr const leaf_type & raw_leaf() const noexcept { return leaf_; }
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
constexpr leaf_type & raw_leaf() noexcept { return leaf_; }
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
const dpf::beaver<false, NodeT, OutputT> & beaver() const noexcept
|
||||
|
|
@ -141,10 +156,11 @@ struct leaf_wrapper<wildcard_value<ConcreteOutputT>, NodeT>
|
|||
|
||||
leaf_wrapper() = delete;
|
||||
leaf_wrapper(leaf_type leaf_share, beaver_type beaver)
|
||||
: leaf_{std::forward<leaf_type>(leaf_share)},
|
||||
: leaf_{leaf_share},
|
||||
beaver_{beaver},
|
||||
output_share_{},
|
||||
leaf_state_{leaf_status::notset}
|
||||
leaf_state_{leaf_status::notset},
|
||||
updating_{false}
|
||||
{ }
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
|
|
@ -157,11 +173,12 @@ struct leaf_wrapper<wildcard_value<ConcreteOutputT>, NodeT>
|
|||
return leaf_;
|
||||
}
|
||||
|
||||
const output_type compute_and_get_blinded_output_share(output_type output_share)
|
||||
output_type compute_and_get_blinded_output_share(output_type output_share)
|
||||
{
|
||||
begin_transition(leaf_status::notset);
|
||||
output_share_ = output_share;
|
||||
auto blinded_output_share = output_share_ + beaver_.output_blind;
|
||||
// Use the leaf group (XOR of IEEE bits for float/double), not `operator+`.
|
||||
auto blinded_output_share = leaf_group_add(output_share_, beaver_.output_blind);
|
||||
leaf_state_ = leaf_status::blinded;
|
||||
return blinded_output_share;
|
||||
}
|
||||
|
|
@ -170,16 +187,28 @@ struct leaf_wrapper<wildcard_value<ConcreteOutputT>, NodeT>
|
|||
/// @tparam Party party index, `0` or `1`
|
||||
/// @tparam Scheme scheme
|
||||
/// @param output_share the `output_share`
|
||||
/// @return the returned `const output_type`
|
||||
/// @return the blinded output share
|
||||
template <std::size_t Party, sharing Scheme>
|
||||
const output_type compute_and_get_blinded_output_share(
|
||||
output_type compute_and_get_blinded_output_share(
|
||||
const secret_share<output_type, Party, Scheme> & output_share)
|
||||
{
|
||||
return compute_and_get_blinded_output_share(
|
||||
output_share.as_additive().raw());
|
||||
// Beaver leaf math is a (2,2) additive absorb. (3,3) and replicated
|
||||
// shares are a different party count; fold them with `add_replicated`.
|
||||
if constexpr (is_two_party_sharing_v<Scheme>)
|
||||
{
|
||||
return compute_and_get_blinded_output_share(
|
||||
output_share.as_additive().raw());
|
||||
}
|
||||
else
|
||||
{
|
||||
static_assert(is_two_party_sharing_v<Scheme>,
|
||||
"wildcard Beaver absorb expects a (2,2) additive or "
|
||||
"subtractive share");
|
||||
return output_type{};
|
||||
}
|
||||
}
|
||||
|
||||
const leaf_type compute_and_get_leaf_share(output_type other_output_share)
|
||||
leaf_type compute_and_get_leaf_share(output_type other_output_share)
|
||||
{
|
||||
begin_transition(leaf_status::blinded);
|
||||
leaf_ = add_leaf<output_type>(leaf_, subtract_leaf<output_type>(
|
||||
|
|
@ -189,10 +218,16 @@ struct leaf_wrapper<wildcard_value<ConcreteOutputT>, NodeT>
|
|||
return leaf_;
|
||||
}
|
||||
|
||||
const leaf_type reconstruct_correction_word(leaf_type other_share)
|
||||
leaf_type reconstruct_correction_word(leaf_type other_share)
|
||||
{
|
||||
begin_transition(leaf_status::waiting);
|
||||
leaf_ = add_leaf<output_type>(leaf_, other_share);
|
||||
if (updating_)
|
||||
{
|
||||
// Opened naked delta; add it onto the previously committed payload.
|
||||
leaf_ = add_leaf<output_type>(committed_, leaf_);
|
||||
updating_ = false;
|
||||
}
|
||||
leaf_state_ = leaf_status::ready;
|
||||
return leaf_;
|
||||
}
|
||||
|
|
@ -211,6 +246,10 @@ struct leaf_wrapper<wildcard_value<ConcreteOutputT>, NodeT>
|
|||
HEDLEY_NO_THROW
|
||||
const leaf_type & raw_leaf() const noexcept { return leaf_; }
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
leaf_type & raw_leaf() noexcept { return leaf_; }
|
||||
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
HEDLEY_NO_THROW
|
||||
const beaver_type & beaver() const noexcept { return beaver_; }
|
||||
|
|
@ -241,6 +280,23 @@ struct leaf_wrapper<wildcard_value<ConcreteOutputT>, NodeT>
|
|||
leaf_state_ = static_cast<leaf_status>(state);
|
||||
}
|
||||
|
||||
/// @brief Re-open a ready leaf so a later assign installs `β' − β`.
|
||||
/// @details Saves the committed correction word and restores the keygen
|
||||
/// Beaver pad (not the zero-payload CW) so the next opening adds
|
||||
/// only a naked delta. Reusing one scale Beaver for two openings
|
||||
/// is not maliciously secure; honest parties that install `β'−β`
|
||||
/// still get a correct leaf.
|
||||
HEDLEY_ALWAYS_INLINE
|
||||
void begin_update()
|
||||
{
|
||||
if (leaf_state_ != leaf_status::ready)
|
||||
throw std::runtime_error("begin_update: leaf is not ready");
|
||||
committed_ = leaf_;
|
||||
leaf_ = beaver_.assign_pad;
|
||||
updating_ = true;
|
||||
leaf_state_ = leaf_status::notset;
|
||||
}
|
||||
|
||||
private:
|
||||
enum class leaf_status : psnip_uint8_t { ready = 0, waiting = 1, computing = 2, blinded = 3, notset = 4 };
|
||||
|
||||
|
|
@ -254,9 +310,11 @@ struct leaf_wrapper<wildcard_value<ConcreteOutputT>, NodeT>
|
|||
}
|
||||
|
||||
leaf_type leaf_;
|
||||
leaf_type committed_{};
|
||||
beaver_type beaver_;
|
||||
output_type output_share_;
|
||||
leaf_status leaf_state_;
|
||||
bool updating_;
|
||||
};
|
||||
|
||||
} // namespace dpf
|
||||
|
|
|
|||
|
|
@ -15,10 +15,15 @@ namespace dpf
|
|||
namespace literals
|
||||
{
|
||||
|
||||
/// @brief Re-export of `N_uN` modint literals.
|
||||
namespace modints{} using namespace dpf::literals::modints;
|
||||
/// @brief Re-export of `N_xN` xint literals.
|
||||
namespace xints{} using namespace dpf::literals::xints;
|
||||
/// @brief Re-export of bitstring literals.
|
||||
namespace bitstrings{} using namespace dpf::literals::bitstrings;
|
||||
/// @brief Re-export of `_twobit`.
|
||||
namespace twobit{} using namespace dpf::literals::twobit;
|
||||
/// @brief Re-export of `_nyble`.
|
||||
namespace nyble{} using namespace dpf::literals::nyble;
|
||||
|
||||
} // namespace literals
|
||||
|
|
|
|||
701
include/dpf/log.hpp
Normal file
701
include/dpf/log.hpp
Normal file
|
|
@ -0,0 +1,701 @@
|
|||
/// @file dpf/log.hpp
|
||||
/// @brief Leveled run log: one `key=value` line per record, written to
|
||||
/// stderr, syslog, or a file.
|
||||
/// @details Nothing is written until `log::configure` runs (normally through
|
||||
/// `app::start_logging`), so library code logs freely and a test or
|
||||
/// embedder that never configures the log stays quiet.
|
||||
///
|
||||
/// A record is built in a local string and written with one
|
||||
/// `write(2)` per sink, so party threads, and party processes that
|
||||
/// append to one file, never interleave inside a line. Every line
|
||||
/// starts with `ts=<UTC, microseconds> lvl=<level> inv=<invocation id>
|
||||
/// pid=<pid> tid=<kernel thread id>`, then `role=<party>` when the
|
||||
/// thread has one (`role_scope`), then `ev=<event>` and the event's
|
||||
/// fields. A value with a space, quote, `=`, backslash, or control
|
||||
/// character is double-quoted with backslash escapes.
|
||||
///
|
||||
/// Seed bytes go through `record::seed`, which prints them as hex, as
|
||||
/// a SHA-256 fingerprint, or not at all, per `settings::seeds`.
|
||||
///
|
||||
/// Levels, from quiet to loud: `silent`, `error`, `warning`, `info`,
|
||||
/// `debug`, `trace` (or `0`-`5`; `warn` and `absurd` are accepted as
|
||||
/// aliases). `DPF_LOG(level, "event").kv("key", value)` builds the
|
||||
/// record only when that level is enabled.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_LOG_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_LOG_HPP__
|
||||
|
||||
#include <array>
|
||||
#include <atomic>
|
||||
#include <cerrno>
|
||||
#include <chrono>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstdio>
|
||||
#include <cstring>
|
||||
#include <ctime>
|
||||
#include <functional>
|
||||
#include <mutex>
|
||||
#include <random>
|
||||
#include <set>
|
||||
#include <stdexcept>
|
||||
#include <string>
|
||||
#include <thread>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
|
||||
#include <fcntl.h>
|
||||
#include <sys/types.h>
|
||||
#include <unistd.h>
|
||||
|
||||
#if defined(__linux__)
|
||||
#include <sys/syscall.h>
|
||||
#endif
|
||||
|
||||
#if defined(__has_include)
|
||||
#if __has_include(<syslog.h>)
|
||||
#include <syslog.h>
|
||||
#define DPF_LOG_HAS_SYSLOG 1
|
||||
#endif
|
||||
#endif
|
||||
#ifndef DPF_LOG_HAS_SYSLOG
|
||||
#define DPF_LOG_HAS_SYSLOG 0
|
||||
#endif
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace log
|
||||
{
|
||||
|
||||
enum class level : unsigned char
|
||||
{
|
||||
silent = 0,
|
||||
error = 1,
|
||||
warning = 2,
|
||||
info = 3,
|
||||
debug = 4,
|
||||
trace = 5
|
||||
};
|
||||
|
||||
inline const char * level_name(level l) noexcept
|
||||
{
|
||||
switch (l)
|
||||
{
|
||||
case level::silent:
|
||||
return "silent";
|
||||
case level::error:
|
||||
return "error";
|
||||
case level::warning:
|
||||
return "warning";
|
||||
case level::info:
|
||||
return "info";
|
||||
case level::debug:
|
||||
return "debug";
|
||||
case level::trace:
|
||||
return "trace";
|
||||
}
|
||||
return "info";
|
||||
}
|
||||
|
||||
/// @brief Parse a level name or digit. Unknown names throw.
|
||||
inline level parse_level(const std::string & s)
|
||||
{
|
||||
if (s == "silent" || s == "none" || s == "off" || s == "0")
|
||||
return level::silent;
|
||||
if (s == "error" || s == "1")
|
||||
return level::error;
|
||||
if (s == "warning" || s == "warn" || s == "2")
|
||||
return level::warning;
|
||||
if (s == "info" || s == "3")
|
||||
return level::info;
|
||||
if (s == "debug" || s == "4")
|
||||
return level::debug;
|
||||
if (s == "trace" || s == "absurd" || s == "5")
|
||||
return level::trace;
|
||||
throw std::invalid_argument("unknown log level: " + s
|
||||
+ " (silent|error|warning|info|debug|trace or 0-5)");
|
||||
}
|
||||
|
||||
/// @brief How seed bytes appear in records.
|
||||
enum class seed_policy : unsigned char
|
||||
{
|
||||
full, ///< the bytes as hex: the run can be replayed from the log
|
||||
hash, ///< a SHA-256 fingerprint: runs can be compared, not replayed
|
||||
off ///< neither; the field reads `withheld`
|
||||
};
|
||||
|
||||
inline const char * seed_policy_name(seed_policy p) noexcept
|
||||
{
|
||||
switch (p)
|
||||
{
|
||||
case seed_policy::full:
|
||||
return "full";
|
||||
case seed_policy::hash:
|
||||
return "hash";
|
||||
case seed_policy::off:
|
||||
return "off";
|
||||
}
|
||||
return "full";
|
||||
}
|
||||
|
||||
inline seed_policy parse_seed_policy(const std::string & s)
|
||||
{
|
||||
if (s == "full" || s == "hex" || s == "1")
|
||||
return seed_policy::full;
|
||||
if (s == "hash" || s == "fingerprint" || s == "sha256")
|
||||
return seed_policy::hash;
|
||||
if (s == "off" || s == "none" || s == "0")
|
||||
return seed_policy::off;
|
||||
throw std::invalid_argument("unknown log_seeds: " + s + " (full|hash|off)");
|
||||
}
|
||||
|
||||
struct settings
|
||||
{
|
||||
level threshold = level::info;
|
||||
/// Comma-separated sinks: `stderr`, `syslog`, `file:PATH`, or `none`.
|
||||
/// One file at most; it is opened for append, so parties may share it.
|
||||
std::string sinks = "stderr";
|
||||
seed_policy seeds = seed_policy::full;
|
||||
/// syslog identity (facility `LOG_USER`).
|
||||
std::string ident = "libdpf";
|
||||
};
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
inline std::atomic<unsigned> threshold{0};
|
||||
inline std::atomic<unsigned char> seeds{0};
|
||||
|
||||
struct sinks_state
|
||||
{
|
||||
std::mutex mu;
|
||||
bool to_stderr = false;
|
||||
bool to_syslog = false;
|
||||
int fd = -1;
|
||||
std::string path;
|
||||
std::string ident;
|
||||
std::set<std::string> said;
|
||||
|
||||
sinks_state() = default;
|
||||
sinks_state(const sinks_state &) = delete;
|
||||
sinks_state & operator=(const sinks_state &) = delete;
|
||||
|
||||
~sinks_state()
|
||||
{
|
||||
threshold.store(0, std::memory_order_relaxed);
|
||||
if (fd >= 0)
|
||||
::close(fd);
|
||||
#if DPF_LOG_HAS_SYSLOG
|
||||
if (to_syslog)
|
||||
::closelog();
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
inline sinks_state & sinks()
|
||||
{
|
||||
static sinks_state s;
|
||||
return s;
|
||||
}
|
||||
|
||||
inline thread_local std::string role_tls;
|
||||
|
||||
inline long kernel_tid() noexcept
|
||||
{
|
||||
#if defined(__linux__) && defined(SYS_gettid)
|
||||
static thread_local const long tid = static_cast<long>(::syscall(SYS_gettid));
|
||||
return tid;
|
||||
#else
|
||||
return static_cast<long>(
|
||||
std::hash<std::thread::id>{}(std::this_thread::get_id()) & 0x7fffffffu);
|
||||
#endif
|
||||
}
|
||||
|
||||
/// @brief `YYYY-MM-DDTHH:MM:SS.uuuuuuZ` for `tp`.
|
||||
inline std::string utc_text(std::chrono::system_clock::time_point tp)
|
||||
{
|
||||
const auto us = std::chrono::duration_cast<std::chrono::microseconds>(
|
||||
tp.time_since_epoch()).count();
|
||||
const std::time_t secs = static_cast<std::time_t>(us / 1000000);
|
||||
const auto frac = static_cast<unsigned>((us % 1000000 + 1000000) % 1000000);
|
||||
std::tm tm{};
|
||||
::gmtime_r(&secs, &tm);
|
||||
char buf[96];
|
||||
std::snprintf(buf, sizeof(buf), "%04d-%02d-%02dT%02d:%02d:%02d.%06uZ",
|
||||
tm.tm_year + 1900, tm.tm_mon + 1, tm.tm_mday, tm.tm_hour, tm.tm_min,
|
||||
tm.tm_sec, frac);
|
||||
return buf;
|
||||
}
|
||||
|
||||
inline void append_value(std::string & out, const char * p, std::size_t n)
|
||||
{
|
||||
bool plain = n != 0;
|
||||
for (std::size_t i = 0; i < n && plain; ++i)
|
||||
{
|
||||
const auto c = static_cast<unsigned char>(p[i]);
|
||||
plain = c > 0x20 && c < 0x7f && c != '"' && c != '=' && c != '\\';
|
||||
}
|
||||
if (plain)
|
||||
{
|
||||
out.append(p, n);
|
||||
return;
|
||||
}
|
||||
out += '"';
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
const auto c = static_cast<unsigned char>(p[i]);
|
||||
switch (c)
|
||||
{
|
||||
case '"':
|
||||
out += "\\\"";
|
||||
break;
|
||||
case '\\':
|
||||
out += "\\\\";
|
||||
break;
|
||||
case '\n':
|
||||
out += "\\n";
|
||||
break;
|
||||
case '\r':
|
||||
out += "\\r";
|
||||
break;
|
||||
case '\t':
|
||||
out += "\\t";
|
||||
break;
|
||||
default:
|
||||
if (c < 0x20 || c == 0x7f)
|
||||
{
|
||||
char esc[5];
|
||||
std::snprintf(esc, sizeof(esc), "\\x%02x", c);
|
||||
out += esc;
|
||||
}
|
||||
else
|
||||
out += static_cast<char>(c);
|
||||
}
|
||||
}
|
||||
out += '"';
|
||||
}
|
||||
|
||||
inline void append_hex(std::string & out, const std::uint8_t * p, std::size_t n)
|
||||
{
|
||||
static const char digits[] = "0123456789abcdef";
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
out += digits[p[i] >> 4];
|
||||
out += digits[p[i] & 0x0f];
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief SHA-256 of `label || data` (FIPS 180-4), for seed fingerprints.
|
||||
inline std::array<std::uint8_t, 32> sha256(const char * label,
|
||||
const std::uint8_t * data, std::size_t n)
|
||||
{
|
||||
static constexpr std::uint32_t k[64] = {
|
||||
0x428a2f98u, 0x71374491u, 0xb5c0fbcfu, 0xe9b5dba5u, 0x3956c25bu,
|
||||
0x59f111f1u, 0x923f82a4u, 0xab1c5ed5u, 0xd807aa98u, 0x12835b01u,
|
||||
0x243185beu, 0x550c7dc3u, 0x72be5d74u, 0x80deb1feu, 0x9bdc06a7u,
|
||||
0xc19bf174u, 0xe49b69c1u, 0xefbe4786u, 0x0fc19dc6u, 0x240ca1ccu,
|
||||
0x2de92c6fu, 0x4a7484aau, 0x5cb0a9dcu, 0x76f988dau, 0x983e5152u,
|
||||
0xa831c66du, 0xb00327c8u, 0xbf597fc7u, 0xc6e00bf3u, 0xd5a79147u,
|
||||
0x06ca6351u, 0x14292967u, 0x27b70a85u, 0x2e1b2138u, 0x4d2c6dfcu,
|
||||
0x53380d13u, 0x650a7354u, 0x766a0abbu, 0x81c2c92eu, 0x92722c85u,
|
||||
0xa2bfe8a1u, 0xa81a664bu, 0xc24b8b70u, 0xc76c51a3u, 0xd192e819u,
|
||||
0xd6990624u, 0xf40e3585u, 0x106aa070u, 0x19a4c116u, 0x1e376c08u,
|
||||
0x2748774cu, 0x34b0bcb5u, 0x391c0cb3u, 0x4ed8aa4au, 0x5b9cca4fu,
|
||||
0x682e6ff3u, 0x748f82eeu, 0x78a5636fu, 0x84c87814u, 0x8cc70208u,
|
||||
0x90befffau, 0xa4506cebu, 0xbef9a3f7u, 0xc67178f2u};
|
||||
std::uint32_t h[8] = {0x6a09e667u, 0xbb67ae85u, 0x3c6ef372u, 0xa54ff53au,
|
||||
0x510e527fu, 0x9b05688cu, 0x1f83d9abu, 0x5be0cd19u};
|
||||
std::string m(label == nullptr ? "" : label);
|
||||
m.append(reinterpret_cast<const char *>(data), n);
|
||||
const std::uint64_t bits = static_cast<std::uint64_t>(m.size()) * 8u;
|
||||
m += static_cast<char>(0x80);
|
||||
while (m.size() % 64 != 56)
|
||||
m += '\0';
|
||||
for (int i = 7; i >= 0; --i)
|
||||
m += static_cast<char>((bits >> (8 * i)) & 0xffu);
|
||||
auto rotr = [](std::uint32_t x, int r) { return (x >> r) | (x << (32 - r)); };
|
||||
for (std::size_t off = 0; off < m.size(); off += 64)
|
||||
{
|
||||
std::uint32_t w[64];
|
||||
for (int i = 0; i < 16; ++i)
|
||||
{
|
||||
const auto * b = reinterpret_cast<const unsigned char *>(m.data() + off + 4 * i);
|
||||
w[i] = (static_cast<std::uint32_t>(b[0]) << 24)
|
||||
| (static_cast<std::uint32_t>(b[1]) << 16)
|
||||
| (static_cast<std::uint32_t>(b[2]) << 8) | static_cast<std::uint32_t>(b[3]);
|
||||
}
|
||||
for (int i = 16; i < 64; ++i)
|
||||
{
|
||||
const std::uint32_t s0 = rotr(w[i - 15], 7) ^ rotr(w[i - 15], 18) ^ (w[i - 15] >> 3);
|
||||
const std::uint32_t s1 = rotr(w[i - 2], 17) ^ rotr(w[i - 2], 19) ^ (w[i - 2] >> 10);
|
||||
w[i] = w[i - 16] + s0 + w[i - 7] + s1;
|
||||
}
|
||||
std::uint32_t a = h[0], b = h[1], c = h[2], d = h[3];
|
||||
std::uint32_t e = h[4], f = h[5], g = h[6], hh = h[7];
|
||||
for (int i = 0; i < 64; ++i)
|
||||
{
|
||||
const std::uint32_t s1 = rotr(e, 6) ^ rotr(e, 11) ^ rotr(e, 25);
|
||||
const std::uint32_t ch = (e & f) ^ (~e & g);
|
||||
const std::uint32_t t1 = hh + s1 + ch + k[i] + w[i];
|
||||
const std::uint32_t s0 = rotr(a, 2) ^ rotr(a, 13) ^ rotr(a, 22);
|
||||
const std::uint32_t mj = (a & b) ^ (a & c) ^ (b & c);
|
||||
const std::uint32_t t2 = s0 + mj;
|
||||
hh = g;
|
||||
g = f;
|
||||
f = e;
|
||||
e = d + t1;
|
||||
d = c;
|
||||
c = b;
|
||||
b = a;
|
||||
a = t1 + t2;
|
||||
}
|
||||
h[0] += a;
|
||||
h[1] += b;
|
||||
h[2] += c;
|
||||
h[3] += d;
|
||||
h[4] += e;
|
||||
h[5] += f;
|
||||
h[6] += g;
|
||||
h[7] += hh;
|
||||
}
|
||||
std::array<std::uint8_t, 32> out{};
|
||||
for (int i = 0; i < 8; ++i)
|
||||
for (int j = 0; j < 4; ++j)
|
||||
out[static_cast<std::size_t>(4 * i + j)] =
|
||||
static_cast<std::uint8_t>((h[i] >> (24 - 8 * j)) & 0xffu);
|
||||
return out;
|
||||
}
|
||||
|
||||
#if DPF_LOG_HAS_SYSLOG
|
||||
inline int syslog_priority(level l) noexcept
|
||||
{
|
||||
switch (l)
|
||||
{
|
||||
case level::error:
|
||||
return LOG_ERR;
|
||||
case level::warning:
|
||||
return LOG_WARNING;
|
||||
case level::info:
|
||||
return LOG_INFO;
|
||||
default:
|
||||
return LOG_DEBUG;
|
||||
}
|
||||
}
|
||||
#endif
|
||||
|
||||
inline void write_all(int fd, const char * p, std::size_t n) noexcept
|
||||
{
|
||||
while (n > 0)
|
||||
{
|
||||
const ssize_t r = ::write(fd, p, n);
|
||||
if (r > 0)
|
||||
{
|
||||
p += r;
|
||||
n -= static_cast<std::size_t>(r);
|
||||
continue;
|
||||
}
|
||||
if (r < 0 && errno == EINTR)
|
||||
continue;
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
inline void emit(level l, std::string & line) noexcept
|
||||
{
|
||||
try
|
||||
{
|
||||
line += '\n';
|
||||
auto & s = sinks();
|
||||
std::lock_guard<std::mutex> lock(s.mu);
|
||||
if (s.to_stderr)
|
||||
write_all(STDERR_FILENO, line.data(), line.size());
|
||||
if (s.fd >= 0)
|
||||
write_all(s.fd, line.data(), line.size());
|
||||
#if DPF_LOG_HAS_SYSLOG
|
||||
if (s.to_syslog)
|
||||
::syslog(syslog_priority(l), "%.*s", static_cast<int>(line.size() - 1),
|
||||
line.data());
|
||||
#else
|
||||
(void)l;
|
||||
#endif
|
||||
}
|
||||
catch (...)
|
||||
{
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/// @brief True when records at `l` are written somewhere.
|
||||
inline bool enabled(level l) noexcept
|
||||
{
|
||||
return l != level::silent
|
||||
&& static_cast<unsigned>(l) <= detail::threshold.load(std::memory_order_relaxed);
|
||||
}
|
||||
|
||||
inline seed_policy seeds() noexcept
|
||||
{
|
||||
return static_cast<seed_policy>(detail::seeds.load(std::memory_order_relaxed));
|
||||
}
|
||||
|
||||
/// @brief Random per-process id stamped on every record (and on CSV rows that
|
||||
/// want to point back at the log). Drawn once from `std::random_device`,
|
||||
/// never from `dpf::uniform_fill`, so it does not perturb seed streams
|
||||
/// or random-byte counts.
|
||||
inline const std::string & invocation_id()
|
||||
{
|
||||
static const std::string id = [] {
|
||||
std::uint64_t v = 0;
|
||||
try
|
||||
{
|
||||
std::random_device rd;
|
||||
v = (static_cast<std::uint64_t>(rd()) << 32) ^ rd();
|
||||
}
|
||||
catch (...)
|
||||
{
|
||||
v = static_cast<std::uint64_t>(
|
||||
std::chrono::system_clock::now().time_since_epoch().count())
|
||||
^ (static_cast<std::uint64_t>(::getpid()) << 40);
|
||||
}
|
||||
std::string out;
|
||||
std::uint8_t b[8];
|
||||
for (int i = 0; i < 8; ++i)
|
||||
b[i] = static_cast<std::uint8_t>((v >> (56 - 8 * i)) & 0xffu);
|
||||
detail::append_hex(out, b, 8);
|
||||
return out;
|
||||
}();
|
||||
return id;
|
||||
}
|
||||
|
||||
/// @brief Apply `cfg`. A bad sink name or unopenable file throws and leaves the
|
||||
/// previous configuration in place. `none` (or no sinks) silences.
|
||||
inline void configure(const settings & cfg)
|
||||
{
|
||||
bool want_stderr = false;
|
||||
bool want_syslog = false;
|
||||
std::string path;
|
||||
std::size_t start = 0;
|
||||
while (start <= cfg.sinks.size())
|
||||
{
|
||||
const auto comma = cfg.sinks.find(',', start);
|
||||
std::string tok = cfg.sinks.substr(start,
|
||||
comma == std::string::npos ? std::string::npos : comma - start);
|
||||
while (!tok.empty() && tok.front() == ' ')
|
||||
tok.erase(tok.begin());
|
||||
while (!tok.empty() && tok.back() == ' ')
|
||||
tok.pop_back();
|
||||
if (tok == "stderr")
|
||||
want_stderr = true;
|
||||
else if (tok == "syslog")
|
||||
want_syslog = true;
|
||||
else if (tok.rfind("file:", 0) == 0 && tok.size() > 5)
|
||||
path = tok.substr(5);
|
||||
else if (!tok.empty() && tok != "none")
|
||||
throw std::invalid_argument("unknown log sink '" + tok
|
||||
+ "' (stderr|syslog|file:PATH|none)");
|
||||
if (comma == std::string::npos)
|
||||
break;
|
||||
start = comma + 1;
|
||||
}
|
||||
#if !DPF_LOG_HAS_SYSLOG
|
||||
if (want_syslog)
|
||||
throw std::invalid_argument("log sink 'syslog' is not available here");
|
||||
#endif
|
||||
(void)invocation_id();
|
||||
auto & s = detail::sinks();
|
||||
std::lock_guard<std::mutex> lock(s.mu);
|
||||
int fd = -1;
|
||||
if (!path.empty())
|
||||
{
|
||||
if (path == s.path && s.fd >= 0)
|
||||
fd = s.fd;
|
||||
else
|
||||
{
|
||||
fd = ::open(path.c_str(), O_WRONLY | O_CREAT | O_APPEND | O_CLOEXEC, 0640);
|
||||
if (fd < 0)
|
||||
throw std::runtime_error("log: cannot open '" + path + "': "
|
||||
+ std::strerror(errno));
|
||||
}
|
||||
}
|
||||
if (s.fd >= 0 && s.fd != fd)
|
||||
::close(s.fd);
|
||||
s.fd = fd;
|
||||
s.path = path;
|
||||
#if DPF_LOG_HAS_SYSLOG
|
||||
if (s.to_syslog && (!want_syslog || s.ident != cfg.ident))
|
||||
{
|
||||
::closelog();
|
||||
s.to_syslog = false;
|
||||
}
|
||||
if (want_syslog && !s.to_syslog)
|
||||
{
|
||||
s.ident = cfg.ident;
|
||||
::openlog(s.ident.c_str(), LOG_PID | LOG_NDELAY, LOG_USER);
|
||||
}
|
||||
#endif
|
||||
s.to_syslog = want_syslog;
|
||||
s.to_stderr = want_stderr;
|
||||
detail::seeds.store(static_cast<unsigned char>(cfg.seeds), std::memory_order_relaxed);
|
||||
const bool any = want_stderr || want_syslog || fd >= 0;
|
||||
detail::threshold.store(any ? static_cast<unsigned>(cfg.threshold) : 0u,
|
||||
std::memory_order_relaxed);
|
||||
}
|
||||
|
||||
/// @brief True the first time `key` is seen in this process (for one-shot
|
||||
/// warnings from code that runs per trial or per cell).
|
||||
inline bool first_time(const std::string & key)
|
||||
{
|
||||
auto & s = detail::sinks();
|
||||
std::lock_guard<std::mutex> lock(s.mu);
|
||||
return s.said.insert(key).second;
|
||||
}
|
||||
|
||||
/// @brief The party role stamped on this thread's records (`p0`, `p1`, `p2`).
|
||||
inline const std::string & role() noexcept { return detail::role_tls; }
|
||||
|
||||
/// @brief Set this thread's role for the scope's lifetime.
|
||||
class role_scope
|
||||
{
|
||||
public:
|
||||
explicit role_scope(std::string role) : prev_(std::move(detail::role_tls))
|
||||
{
|
||||
detail::role_tls = std::move(role);
|
||||
}
|
||||
~role_scope() { detail::role_tls = std::move(prev_); }
|
||||
role_scope(const role_scope &) = delete;
|
||||
role_scope & operator=(const role_scope &) = delete;
|
||||
|
||||
private:
|
||||
std::string prev_;
|
||||
};
|
||||
|
||||
/// @brief One log line. Written when it goes out of scope.
|
||||
class record
|
||||
{
|
||||
public:
|
||||
record(level l, const char * event) : level_(l)
|
||||
{
|
||||
line_.reserve(256);
|
||||
line_ += "ts=";
|
||||
line_ += detail::utc_text(std::chrono::system_clock::now());
|
||||
line_ += " lvl=";
|
||||
line_ += level_name(l);
|
||||
line_ += " inv=";
|
||||
line_ += invocation_id();
|
||||
line_ += " pid=";
|
||||
line_ += std::to_string(::getpid());
|
||||
line_ += " tid=";
|
||||
line_ += std::to_string(detail::kernel_tid());
|
||||
if (!detail::role_tls.empty())
|
||||
{
|
||||
line_ += " role=";
|
||||
detail::append_value(line_, detail::role_tls.data(), detail::role_tls.size());
|
||||
}
|
||||
line_ += " ev=";
|
||||
const char * ev = event == nullptr ? "event" : event;
|
||||
detail::append_value(line_, ev, std::strlen(ev));
|
||||
}
|
||||
|
||||
record(const record &) = delete;
|
||||
record & operator=(const record &) = delete;
|
||||
|
||||
~record() { detail::emit(level_, line_); }
|
||||
|
||||
record & kv(const char * key, const std::string & v)
|
||||
{
|
||||
key_(key);
|
||||
detail::append_value(line_, v.data(), v.size());
|
||||
return *this;
|
||||
}
|
||||
|
||||
record & kv(const char * key, const char * v)
|
||||
{
|
||||
key_(key);
|
||||
const char * p = v == nullptr ? "" : v;
|
||||
detail::append_value(line_, p, std::strlen(p));
|
||||
return *this;
|
||||
}
|
||||
|
||||
record & kv(const char * key, bool v)
|
||||
{
|
||||
key_(key);
|
||||
line_ += v ? '1' : '0';
|
||||
return *this;
|
||||
}
|
||||
|
||||
template <typename T,
|
||||
std::enable_if_t<std::is_integral_v<T> && !std::is_same_v<T, bool>, int> = 0>
|
||||
record & kv(const char * key, T v)
|
||||
{
|
||||
key_(key);
|
||||
line_ += std::to_string(v);
|
||||
return *this;
|
||||
}
|
||||
|
||||
template <typename T, std::enable_if_t<std::is_floating_point_v<T>, int> = 0>
|
||||
record & kv(const char * key, T v)
|
||||
{
|
||||
key_(key);
|
||||
char buf[32];
|
||||
std::snprintf(buf, sizeof(buf), "%.6g", static_cast<double>(v));
|
||||
line_ += buf;
|
||||
return *this;
|
||||
}
|
||||
|
||||
record & hex(const char * key, const void * bytes, std::size_t n)
|
||||
{
|
||||
key_(key);
|
||||
if (n == 0 || bytes == nullptr)
|
||||
line_ += "\"\"";
|
||||
else
|
||||
detail::append_hex(line_, static_cast<const std::uint8_t *>(bytes), n);
|
||||
return *this;
|
||||
}
|
||||
|
||||
/// @brief Seed bytes under the configured `seed_policy`.
|
||||
record & seed(const char * key, const void * bytes, std::size_t n)
|
||||
{
|
||||
switch (seeds())
|
||||
{
|
||||
case seed_policy::full:
|
||||
return hex(key, bytes, n);
|
||||
case seed_policy::hash:
|
||||
{
|
||||
const auto d = detail::sha256("libdpf seed fingerprint",
|
||||
static_cast<const std::uint8_t *>(bytes), bytes == nullptr ? 0 : n);
|
||||
key_(key);
|
||||
line_ += "sha256:";
|
||||
detail::append_hex(line_, d.data(), 8);
|
||||
return *this;
|
||||
}
|
||||
case seed_policy::off:
|
||||
break;
|
||||
}
|
||||
key_(key);
|
||||
line_ += "withheld";
|
||||
return *this;
|
||||
}
|
||||
|
||||
private:
|
||||
void key_(const char * key)
|
||||
{
|
||||
line_ += ' ';
|
||||
line_ += key == nullptr ? "field" : key;
|
||||
line_ += '=';
|
||||
}
|
||||
|
||||
level level_;
|
||||
std::string line_;
|
||||
};
|
||||
|
||||
} // namespace log
|
||||
} // namespace dpf
|
||||
|
||||
/// @brief `DPF_LOG(info, "event").kv("key", value)`: the record and its
|
||||
/// arguments are only evaluated when `info` is enabled. The single-pass
|
||||
/// `for` has no `else` to pair with an enclosing unbraced `if`.
|
||||
#define DPF_LOG(LEVEL, EVENT) \
|
||||
for (bool dpf_log_on_ = ::dpf::log::enabled(::dpf::log::level::LEVEL); \
|
||||
dpf_log_on_; dpf_log_on_ = false) \
|
||||
::dpf::log::record(::dpf::log::level::LEVEL, EVENT)
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_LOG_HPP__
|
||||
173
include/dpf/matmul.hpp
Normal file
173
include/dpf/matmul.hpp
Normal file
|
|
@ -0,0 +1,173 @@
|
|||
/// @file dpf/matmul.hpp
|
||||
/// @brief Matrix Beaver triples and online matmul / im2col convolution.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_MATMUL_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_MATMUL_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <stdexcept>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/random.hpp"
|
||||
#include "dpf/share_vec.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace matmul
|
||||
{
|
||||
|
||||
template <typename Ring>
|
||||
struct dims
|
||||
{
|
||||
std::size_t m = 0;
|
||||
std::size_t k = 0;
|
||||
std::size_t n = 0;
|
||||
};
|
||||
|
||||
template <typename Ring>
|
||||
struct matrix_triple
|
||||
{
|
||||
dims<Ring> d{};
|
||||
std::vector<Ring> a0, a1; ///< m×k
|
||||
std::vector<Ring> b0, b1; ///< k×n
|
||||
std::vector<Ring> c0, c1; ///< m×n = AB
|
||||
};
|
||||
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::vector<Ring> clear_mul(const std::vector<Ring> & a,
|
||||
const std::vector<Ring> & b, dims<Ring> d)
|
||||
{
|
||||
if (a.size() != d.m * d.k || b.size() != d.k * d.n)
|
||||
throw std::invalid_argument("matmul clear size");
|
||||
std::vector<Ring> c(d.m * d.n, Ring{});
|
||||
for (std::size_t i = 0; i < d.m; ++i)
|
||||
for (std::size_t j = 0; j < d.n; ++j)
|
||||
{
|
||||
Ring acc{};
|
||||
for (std::size_t t = 0; t < d.k; ++t)
|
||||
acc = static_cast<Ring>(
|
||||
acc + a[i * d.k + t] * b[t * d.n + j]);
|
||||
c[i * d.n + j] = acc;
|
||||
}
|
||||
return c;
|
||||
}
|
||||
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
matrix_triple<Ring> sample_triple(dims<Ring> d)
|
||||
{
|
||||
matrix_triple<Ring> t;
|
||||
t.d = d;
|
||||
t.a0.resize(d.m * d.k);
|
||||
t.a1.resize(d.m * d.k);
|
||||
t.b0.resize(d.k * d.n);
|
||||
t.b1.resize(d.k * d.n);
|
||||
std::vector<Ring> a(d.m * d.k), b(d.k * d.n);
|
||||
for (std::size_t i = 0; i < a.size(); ++i)
|
||||
{
|
||||
a[i] = dpf::uniform_sample<Ring>();
|
||||
t.a0[i] = dpf::uniform_sample<Ring>();
|
||||
t.a1[i] = static_cast<Ring>(a[i] - t.a0[i]);
|
||||
}
|
||||
for (std::size_t i = 0; i < b.size(); ++i)
|
||||
{
|
||||
b[i] = dpf::uniform_sample<Ring>();
|
||||
t.b0[i] = dpf::uniform_sample<Ring>();
|
||||
t.b1[i] = static_cast<Ring>(b[i] - t.b0[i]);
|
||||
}
|
||||
auto c = clear_mul(a, b, d);
|
||||
t.c0.resize(c.size());
|
||||
t.c1.resize(c.size());
|
||||
for (std::size_t i = 0; i < c.size(); ++i)
|
||||
{
|
||||
t.c0[i] = dpf::uniform_sample<Ring>();
|
||||
t.c1[i] = static_cast<Ring>(c[i] - t.c0[i]);
|
||||
}
|
||||
return t;
|
||||
}
|
||||
|
||||
/// @brief Online: open `D = X-A`, `E = Y-B`; `Z = C + D B + A E + D E` (p0 adds DE).
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::vector<Ring> online(const std::vector<Ring> & x_share,
|
||||
const std::vector<Ring> & y_share, const std::vector<Ring> & a_share,
|
||||
const std::vector<Ring> & b_share, const std::vector<Ring> & c_share,
|
||||
const std::vector<Ring> & d_open, const std::vector<Ring> & e_open,
|
||||
dims<Ring> d, unsigned party)
|
||||
{
|
||||
if (x_share.size() != d.m * d.k || y_share.size() != d.k * d.n)
|
||||
throw std::invalid_argument("matmul online size");
|
||||
auto db = clear_mul(d_open, b_share, d);
|
||||
// A·E : a is m×k, e is k×n
|
||||
auto ae = clear_mul(a_share, e_open, d);
|
||||
std::vector<Ring> z(d.m * d.n);
|
||||
for (std::size_t i = 0; i < z.size(); ++i)
|
||||
z[i] = static_cast<Ring>(c_share[i] + db[i] + ae[i]);
|
||||
if (party == 0)
|
||||
{
|
||||
auto de = clear_mul(d_open, e_open, d);
|
||||
for (std::size_t i = 0; i < z.size(); ++i)
|
||||
z[i] = static_cast<Ring>(z[i] + de[i]);
|
||||
}
|
||||
(void)x_share;
|
||||
(void)y_share;
|
||||
return z;
|
||||
}
|
||||
|
||||
/// @brief Full clear+share test: return shares of X Y.
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::pair<std::vector<Ring>, std::vector<Ring>> mul_shared(
|
||||
const std::vector<Ring> & x, const std::vector<Ring> & y, dims<Ring> d)
|
||||
{
|
||||
auto t = sample_triple<Ring>(d);
|
||||
std::vector<Ring> x0(x.size()), x1(x.size()), y0(y.size()), y1(y.size());
|
||||
for (std::size_t i = 0; i < x.size(); ++i)
|
||||
{
|
||||
x0[i] = dpf::uniform_sample<Ring>();
|
||||
x1[i] = static_cast<Ring>(x[i] - x0[i]);
|
||||
}
|
||||
for (std::size_t i = 0; i < y.size(); ++i)
|
||||
{
|
||||
y0[i] = dpf::uniform_sample<Ring>();
|
||||
y1[i] = static_cast<Ring>(y[i] - y0[i]);
|
||||
}
|
||||
std::vector<Ring> d_open(d.m * d.k), e_open(d.k * d.n);
|
||||
for (std::size_t i = 0; i < d_open.size(); ++i)
|
||||
d_open[i] = static_cast<Ring>((x0[i] + x1[i]) - (t.a0[i] + t.a1[i]));
|
||||
for (std::size_t i = 0; i < e_open.size(); ++i)
|
||||
e_open[i] = static_cast<Ring>((y0[i] + y1[i]) - (t.b0[i] + t.b1[i]));
|
||||
auto z0 = online(x0, y0, t.a0, t.b0, t.c0, d_open, e_open, d, 0);
|
||||
auto z1 = online(x1, y1, t.a1, t.b1, t.c1, d_open, e_open, d, 1);
|
||||
return {std::move(z0), std::move(z1)};
|
||||
}
|
||||
|
||||
/// @brief im2col: flatten `h×w` patches of size `kh×kw` with stride 1 into rows.
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::vector<Ring> im2col(const std::vector<Ring> & img, std::size_t h,
|
||||
std::size_t w, std::size_t kh, std::size_t kw)
|
||||
{
|
||||
if (img.size() != h * w || kh > h || kw > w)
|
||||
throw std::invalid_argument("im2col");
|
||||
const std::size_t out_h = h - kh + 1;
|
||||
const std::size_t out_w = w - kw + 1;
|
||||
std::vector<Ring> col(out_h * out_w * kh * kw);
|
||||
std::size_t row = 0;
|
||||
for (std::size_t i = 0; i < out_h; ++i)
|
||||
for (std::size_t j = 0; j < out_w; ++j, ++row)
|
||||
for (std::size_t u = 0; u < kh; ++u)
|
||||
for (std::size_t v = 0; v < kw; ++v)
|
||||
col[row * (kh * kw) + u * kw + v] =
|
||||
img[(i + u) * w + (j + v)];
|
||||
return col;
|
||||
}
|
||||
|
||||
} // namespace matmul
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_MATMUL_HPP__
|
||||
91
include/dpf/mesh_apps.hpp
Normal file
91
include/dpf/mesh_apps.hpp
Normal file
|
|
@ -0,0 +1,91 @@
|
|||
/// @file dpf/mesh_apps.hpp
|
||||
/// @brief Micro-plan builders for PIRsona and hushmap on the edge mesh.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_MESH_APPS_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_MESH_APPS_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <memory>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/compose.hpp"
|
||||
#include "dpf/pad_graphs.hpp"
|
||||
#include "dpf/protocol.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace protocol
|
||||
{
|
||||
|
||||
/// @brief PIRsona: BitMore-shaped star fetch (L bit keys → 2^L servers).
|
||||
/// @details Returns upload+answer rounds for `n_servers = 1u << L`.
|
||||
inline std::vector<schedule_round> pirsona_bitmore_fetch(std::size_t L,
|
||||
std::size_t seed_bytes, std::size_t answer_bytes,
|
||||
const std::shared_ptr<std::vector<std::vector<std::uint8_t>>> & seeds,
|
||||
const std::shared_ptr<std::vector<std::vector<std::uint8_t>>> & answers)
|
||||
{
|
||||
const std::size_t n = std::size_t{1} << L;
|
||||
return star_upload_answer_graph(n, seed_bytes * L, answer_bytes, seeds,
|
||||
answers);
|
||||
}
|
||||
|
||||
/// @brief PIRsona one gradient-descent update: Du-Atallah mul graph (stub body).
|
||||
inline std::vector<schedule_round> pirsona_gd_update_graph(
|
||||
const std::shared_ptr<std::vector<std::uint8_t>> & tape)
|
||||
{
|
||||
return du_atallah_mul_graph(tape, /*slot_bytes=*/32);
|
||||
}
|
||||
|
||||
/// @brief Hushmap KHM ADD online skeleton: Beaver open + masked open rounds.
|
||||
/// @details Matches the barrier count of MPC_PROTOCOL-v3 Phase 1 (2 peer opens)
|
||||
/// after a dealer tape of `n_layers` triples has been spliced.
|
||||
inline std::vector<schedule_round> hushmap_add_online_graph(
|
||||
std::size_t n_layers, std::size_t open_bytes = 8)
|
||||
{
|
||||
(void)n_layers;
|
||||
std::vector<schedule_round> rounds(2);
|
||||
for (std::size_t r = 0; r < 2; ++r)
|
||||
{
|
||||
rounds[r].slot_bytes = open_bytes;
|
||||
rounds[r].edge = edge_peer;
|
||||
rounds[r].channel = edge_channel::peer;
|
||||
rounds[r].recv = receive_rule::domain_open;
|
||||
rounds[r].sink_round = static_cast<std::uint16_t>(r);
|
||||
rounds[r].produce = [open_bytes, r](std::size_t, const std::uint8_t *,
|
||||
std::size_t, std::uint8_t * out) {
|
||||
if (out != nullptr)
|
||||
std::memset(out, static_cast<int>(0x10 + r), open_bytes);
|
||||
};
|
||||
}
|
||||
return rounds;
|
||||
}
|
||||
|
||||
/// @brief Full hushmap ADD schedule: dealer tape then online opens.
|
||||
inline std::vector<schedule_round> hushmap_add_schedule(std::size_t n_layers,
|
||||
const std::shared_ptr<std::vector<std::uint8_t>> & tape,
|
||||
std::size_t triple_bytes = 24, std::size_t open_bytes = 8)
|
||||
{
|
||||
auto pads = dealer_tape_graph(n_layers, triple_bytes, tape);
|
||||
auto online = hushmap_add_online_graph(n_layers, open_bytes);
|
||||
// Online peer sink_rounds start at 0 on a peer-only sink after dealer
|
||||
// edges are separate — keep sink_round as assigned.
|
||||
return splice_rounds(std::move(pads), std::move(online));
|
||||
}
|
||||
|
||||
/// @brief Composer plan for a 2-server keyword-PIR-shaped client_servers flow.
|
||||
/// @details Prefer `n_server_pir_plan` / `keyword_pir_compose_plan` in
|
||||
/// [app_plans.hpp](@ref dpf/app_plans.hpp) for depth-derived sizes.
|
||||
inline plan keyword_pir_plan(std::size_t party, std::size_t query_bytes,
|
||||
std::size_t answer_bytes)
|
||||
{
|
||||
composer c(party);
|
||||
auto q = c.client_servers(2, query_bytes, answer_bytes);
|
||||
(void)q;
|
||||
return c.default_plan();
|
||||
}
|
||||
|
||||
} // namespace protocol
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_MESH_APPS_HPP__
|
||||
|
|
@ -34,6 +34,10 @@ namespace dpf
|
|||
|
||||
/// @brief represents an unsigned integer modulo `2^Nbits` for small values of `Nbits`
|
||||
/// @tparam Nbits width in bits
|
||||
/// @note Domain widths 1, 2, and 4 are this type or `dpf::xint`, not
|
||||
/// `dpf::bit` / `dpf::twobit` / `dpf::nyble`.
|
||||
/// @see dpf::xint
|
||||
/// @see dpf::modints::modintN_t
|
||||
template <std::size_t Nbits>
|
||||
class modint
|
||||
{
|
||||
|
|
@ -831,6 +835,12 @@ struct mod_pow_2<dpf::modint<Nbits>>
|
|||
namespace modints
|
||||
{
|
||||
|
||||
/// @name modintN_t
|
||||
/// @{
|
||||
|
||||
/// @brief `modintN_t` is `dpf::modint<N>` for N from 1 through 256.
|
||||
/// @see dpf::modint
|
||||
/// @see dpf::xint
|
||||
// 1--9
|
||||
using modint1_t = dpf::modint<1>;
|
||||
using modint2_t = dpf::modint<2>;
|
||||
|
|
@ -1114,6 +1124,8 @@ using modint254_t = dpf::modint<254>;
|
|||
using modint255_t = dpf::modint<255>;
|
||||
using modint256_t = dpf::modint<256>;
|
||||
|
||||
/// @}
|
||||
|
||||
namespace literals = dpf::literals::modints;
|
||||
|
||||
} // namespace modints
|
||||
|
|
@ -1124,6 +1136,11 @@ namespace literals
|
|||
namespace modints
|
||||
{
|
||||
|
||||
/// @name modint literals `N_uN`
|
||||
/// @{
|
||||
|
||||
/// @brief Decimal literal for `dpf::modint<N>`. Widths through 64 take an integer; wider widths take a digit string.
|
||||
/// @see dpf::modints::modintN_t
|
||||
// 1--9
|
||||
constexpr static auto operator "" _u1(unsigned long long int x) { return dpf::modints::modint1_t{static_cast<psnip_uint8_t>(x)}; }
|
||||
constexpr static auto operator "" _u2(unsigned long long int x) { return dpf::modints::modint2_t{static_cast<psnip_uint8_t>(x)}; }
|
||||
|
|
@ -1408,6 +1425,8 @@ template <char ...digits> constexpr static auto operator "" _u254() { utils::con
|
|||
template <char ...digits> constexpr static auto operator "" _u255() { utils::constexpr_maybe_throw<std::runtime_error>(!(std::isdigit(digits) && ...), "invalid char"); uint256_t x{0}; ((x = x * 10 + (digits - '0')), ...); return dpf::modints::modint255_t{x}; }
|
||||
template <char ...digits> constexpr static auto operator "" _u256() { utils::constexpr_maybe_throw<std::runtime_error>(!(std::isdigit(digits) && ...), "invalid char"); uint256_t x{0}; ((x = x * 10 + (digits - '0')), ...); return dpf::modints::modint256_t{x}; }
|
||||
|
||||
/// @}
|
||||
|
||||
} // namespace modints
|
||||
|
||||
} // namespace literals
|
||||
|
|
|
|||
|
|
@ -1,10 +1,11 @@
|
|||
/// @file dpf/multipoint.hpp
|
||||
/// @brief Cuckoo-packed multi-point DPF and verifiable multi-point DPF.
|
||||
/// @details Packs t distinct points into m ≈ O(t) buckets (de Castro–
|
||||
/// Polychroniadou, EUROCRYPT 2022, §4). Each bucket is an ordinary
|
||||
/// Polychroniadou, EUROCRYPT 2022, §4, ePrint 2021/580). Each bucket is an ordinary
|
||||
/// point key on a smaller domain — `dpf::verifiable` selects VDPF
|
||||
/// buckets. Evaluation probes κ = 3 buckets and sums the shares.
|
||||
/// A batched proof is one 2λ token.
|
||||
/// @note Following that section: κ = 3 cuckoo hashes, one point key per bucket.
|
||||
/// @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.
|
||||
|
|
@ -32,6 +33,7 @@
|
|||
#include "dpf/prg_aes.hpp"
|
||||
#include "dpf/random.hpp"
|
||||
#include "dpf/secret_share.hpp"
|
||||
#include "dpf/uint256_t.hpp"
|
||||
#include "dpf/verifiable.hpp"
|
||||
|
||||
namespace dpf
|
||||
|
|
@ -45,6 +47,26 @@ struct multipoint_params
|
|||
int retries = 8;
|
||||
};
|
||||
|
||||
/// @brief 512-bit word for the cuckoo PRP. Holds `3·2^b` for every input
|
||||
/// width this library can form a point key on (up to 256 bits).
|
||||
struct mpf_word
|
||||
{
|
||||
uint256_t lo{};
|
||||
uint256_t hi{};
|
||||
|
||||
friend bool operator==(mpf_word a, mpf_word b) noexcept
|
||||
{
|
||||
return a.lo == b.lo && a.hi == b.hi;
|
||||
}
|
||||
|
||||
friend bool operator<(mpf_word a, mpf_word b) noexcept
|
||||
{
|
||||
if (a.hi != b.hi)
|
||||
return a.hi < b.hi;
|
||||
return a.lo < b.lo;
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
struct is_multipoint_key : std::false_type
|
||||
{
|
||||
|
|
@ -68,8 +90,8 @@ struct multipoint_key
|
|||
using share_type = subtractive_share<OutputT, Party>;
|
||||
|
||||
simde__m128i sigma{};
|
||||
std::uint32_t bucket_count = 0;
|
||||
std::uint64_t bucket_domain = 0;
|
||||
std::uint64_t bucket_count = 0;
|
||||
mpf_word bucket_domain{};
|
||||
std::vector<party_key<Party, BucketKey>> buckets{};
|
||||
};
|
||||
|
||||
|
|
@ -98,15 +120,237 @@ struct prp_walk_error : std::runtime_error
|
|||
|
||||
struct located
|
||||
{
|
||||
std::uint32_t bucket = 0;
|
||||
std::uint64_t index = 0;
|
||||
std::uint64_t bucket = 0;
|
||||
mpf_word index{};
|
||||
};
|
||||
|
||||
using wide = unsigned __int128;
|
||||
|
||||
inline wide domain_size(std::size_t bits)
|
||||
inline mpf_word word_add(mpf_word a, mpf_word b)
|
||||
{
|
||||
return wide{1} << bits;
|
||||
mpf_word r;
|
||||
r.lo = a.lo + b.lo;
|
||||
r.hi = a.hi + b.hi;
|
||||
if (r.lo < a.lo)
|
||||
r.hi = r.hi + uint256_t{1};
|
||||
return r;
|
||||
}
|
||||
|
||||
inline mpf_word word_sub(mpf_word a, mpf_word b)
|
||||
{
|
||||
mpf_word r;
|
||||
r.lo = a.lo - b.lo;
|
||||
r.hi = a.hi - b.hi;
|
||||
if (a.lo < b.lo)
|
||||
r.hi = r.hi - uint256_t{1};
|
||||
return r;
|
||||
}
|
||||
|
||||
inline mpf_word word_shl(mpf_word a, unsigned shift)
|
||||
{
|
||||
if (shift == 0)
|
||||
return a;
|
||||
if (shift >= 512)
|
||||
return {};
|
||||
if (shift >= 256)
|
||||
{
|
||||
mpf_word r;
|
||||
r.hi = a.lo << (shift - 256);
|
||||
return r;
|
||||
}
|
||||
mpf_word r;
|
||||
r.lo = a.lo << shift;
|
||||
r.hi = (a.hi << shift) | (a.lo >> (256 - shift));
|
||||
return r;
|
||||
}
|
||||
|
||||
inline mpf_word word_shr(mpf_word a, unsigned shift)
|
||||
{
|
||||
if (shift == 0)
|
||||
return a;
|
||||
if (shift >= 512)
|
||||
return {};
|
||||
if (shift >= 256)
|
||||
{
|
||||
mpf_word r;
|
||||
r.lo = a.hi >> (shift - 256);
|
||||
return r;
|
||||
}
|
||||
mpf_word r;
|
||||
r.hi = a.hi >> shift;
|
||||
r.lo = (a.lo >> shift) | (a.hi << (256 - shift));
|
||||
return r;
|
||||
}
|
||||
|
||||
inline mpf_word word_or(mpf_word a, mpf_word b)
|
||||
{
|
||||
a.lo = a.lo | b.lo;
|
||||
a.hi = a.hi | b.hi;
|
||||
return a;
|
||||
}
|
||||
|
||||
inline mpf_word word_and(mpf_word a, mpf_word b)
|
||||
{
|
||||
a.lo = a.lo & b.lo;
|
||||
a.hi = a.hi & b.hi;
|
||||
return a;
|
||||
}
|
||||
|
||||
inline bool word_bit(mpf_word a, unsigned bit)
|
||||
{
|
||||
if (bit >= 512)
|
||||
return false;
|
||||
if (bit >= 256)
|
||||
return static_cast<bool>((a.hi >> (bit - 256)) & uint256_t{1});
|
||||
return static_cast<bool>((a.lo >> bit) & uint256_t{1});
|
||||
}
|
||||
|
||||
inline int word_bit_length(mpf_word a)
|
||||
{
|
||||
for (int i = 255; i >= 0; --i)
|
||||
{
|
||||
if (static_cast<bool>((a.hi >> i) & uint256_t{1}))
|
||||
return i + 1 + 256;
|
||||
}
|
||||
for (int i = 255; i >= 0; --i)
|
||||
{
|
||||
if (static_cast<bool>((a.lo >> i) & uint256_t{1}))
|
||||
return i + 1;
|
||||
}
|
||||
return 0;
|
||||
}
|
||||
|
||||
inline mpf_word word_mul_small(mpf_word a, std::uint64_t k)
|
||||
{
|
||||
mpf_word r{};
|
||||
while (k != 0)
|
||||
{
|
||||
if ((k & 1u) != 0)
|
||||
r = word_add(r, a);
|
||||
a = word_shl(a, 1);
|
||||
k >>= 1;
|
||||
}
|
||||
return r;
|
||||
}
|
||||
|
||||
inline std::pair<mpf_word, mpf_word> word_divmod(mpf_word num, mpf_word den)
|
||||
{
|
||||
if (den == mpf_word{})
|
||||
throw std::invalid_argument("multipoint division by zero");
|
||||
mpf_word q{};
|
||||
mpf_word r{};
|
||||
const int top = word_bit_length(num);
|
||||
for (int i = top - 1; i >= 0; --i)
|
||||
{
|
||||
r = word_shl(r, 1);
|
||||
if (word_bit(num, static_cast<unsigned>(i)))
|
||||
r = word_add(r, mpf_word{uint256_t{1}, uint256_t{0}});
|
||||
if (!(r < den))
|
||||
{
|
||||
r = word_sub(r, den);
|
||||
mpf_word bit{};
|
||||
if (i >= 256)
|
||||
bit.hi = uint256_t{1} << static_cast<unsigned>(i - 256);
|
||||
else
|
||||
bit.lo = uint256_t{1} << static_cast<unsigned>(i);
|
||||
q = word_or(q, bit);
|
||||
}
|
||||
}
|
||||
return {q, r};
|
||||
}
|
||||
|
||||
inline mpf_word domain_size(std::size_t bits)
|
||||
{
|
||||
mpf_word r{};
|
||||
if (bits >= 512)
|
||||
throw std::invalid_argument("multipoint domain shift is out of range");
|
||||
if (bits >= 256)
|
||||
r.hi = uint256_t{1} << (bits - 256);
|
||||
else if (bits > 0)
|
||||
r.lo = uint256_t{1} << bits;
|
||||
return r;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
mpf_word to_word(T x)
|
||||
{
|
||||
constexpr std::size_t bits = utils::bitlength_of_v<T>;
|
||||
mpf_word w{};
|
||||
if constexpr (bits > 128)
|
||||
{
|
||||
w.lo = static_cast<uint256_t>(x);
|
||||
}
|
||||
else if constexpr (bits > 64)
|
||||
{
|
||||
uint128_t low{};
|
||||
std::memcpy(&low, &x, sizeof(T));
|
||||
w.lo = uint256_t{low};
|
||||
}
|
||||
else
|
||||
{
|
||||
w.lo = uint256_t{static_cast<std::uint64_t>(x)};
|
||||
}
|
||||
return w;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
T from_word(mpf_word w)
|
||||
{
|
||||
constexpr std::size_t bits = utils::bitlength_of_v<T>;
|
||||
if constexpr (bits > 128)
|
||||
{
|
||||
return static_cast<T>(w.lo);
|
||||
}
|
||||
else if constexpr (bits > 64)
|
||||
{
|
||||
const uint128_t low = static_cast<uint128_t>(w.lo);
|
||||
T out{};
|
||||
std::memcpy(&out, &low, sizeof(T));
|
||||
return out;
|
||||
}
|
||||
else
|
||||
{
|
||||
return static_cast<T>(static_cast<std::uint64_t>(w.lo));
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Low `half` bits set, as a 512-bit mask. `half <= 0` is zero.
|
||||
inline mpf_word low_mask(int half)
|
||||
{
|
||||
if (half <= 0)
|
||||
return {};
|
||||
if (half >= 512)
|
||||
{
|
||||
mpf_word all;
|
||||
all.lo = ~uint256_t{0};
|
||||
all.hi = ~uint256_t{0};
|
||||
return all;
|
||||
}
|
||||
return word_sub(domain_size(static_cast<std::size_t>(half)),
|
||||
mpf_word{uint256_t{1}, uint256_t{0}});
|
||||
}
|
||||
|
||||
inline mpf_word aes_prf(simde__m128i seed, mpf_word right, int round)
|
||||
{
|
||||
alignas(16) unsigned char raw[32]{};
|
||||
std::memcpy(raw, &right.lo, sizeof(right.lo));
|
||||
alignas(16) simde__m128i block0;
|
||||
alignas(16) simde__m128i block1;
|
||||
std::memcpy(&block0, raw, 16);
|
||||
std::memcpy(&block1, raw + 16, 16);
|
||||
block0 = simde_mm_xor_si128(block0, seed);
|
||||
block0 = simde_mm_xor_si128(block0, simde_mm_set_epi32(0, 0, 0, round + 1));
|
||||
const auto out0 = prg::aes128::eval(block0,
|
||||
static_cast<psnip_uint32_t>(round + 1));
|
||||
block1 = simde_mm_xor_si128(block1, seed);
|
||||
block1 = simde_mm_xor_si128(block1,
|
||||
simde_mm_set_epi32(0, 0, 0, round + 0x11));
|
||||
const auto out1 = prg::aes128::eval(block1,
|
||||
static_cast<psnip_uint32_t>(round + 0x21));
|
||||
alignas(16) unsigned char packed[32];
|
||||
std::memcpy(packed, &out0, 16);
|
||||
std::memcpy(packed + 16, &out1, 16);
|
||||
mpf_word f{};
|
||||
std::memcpy(&f.lo, packed, sizeof(f.lo));
|
||||
return f;
|
||||
}
|
||||
|
||||
/// @brief 4-round Feistel on the next power-of-two square, then cycle-walk
|
||||
|
|
@ -117,75 +361,68 @@ inline wide domain_size(std::size_t bits)
|
|||
/// @return the permuted value in `[0, domain)`
|
||||
/// @throws std::invalid_argument if `x` is outside the domain
|
||||
/// @throws prp_walk_error if the cycle walk exceeds its bound
|
||||
inline wide permute(simde__m128i seed, wide x, wide domain)
|
||||
inline mpf_word permute(simde__m128i seed, mpf_word x, mpf_word domain)
|
||||
{
|
||||
if (domain <= 1)
|
||||
return 0;
|
||||
if (x >= domain)
|
||||
const mpf_word one{uint256_t{1}, uint256_t{0}};
|
||||
if (!(one < domain))
|
||||
return {};
|
||||
if (!(x < domain))
|
||||
throw std::invalid_argument("multipoint PRP input is outside the domain");
|
||||
|
||||
int bits = 0;
|
||||
for (wide v = domain - 1; v > 0; v >>= 1)
|
||||
++bits;
|
||||
const int bits = word_bit_length(word_sub(domain, one));
|
||||
const int half = (bits + 1) / 2;
|
||||
const wide mask = (half >= 128)
|
||||
? ~wide{0}
|
||||
: (wide{1} << half) - 1;
|
||||
const mpf_word mask = low_mask(half);
|
||||
|
||||
wide val = x;
|
||||
mpf_word val = x;
|
||||
for (int guard = 0; guard < 128; ++guard)
|
||||
{
|
||||
unsigned __int128 left = (val >> half) & mask;
|
||||
unsigned __int128 right = val & mask;
|
||||
mpf_word left = word_and(word_shr(val, static_cast<unsigned>(half)), mask);
|
||||
mpf_word right = word_and(val, mask);
|
||||
for (int round = 0; round < 4; ++round)
|
||||
{
|
||||
alignas(16) std::uint64_t lanes[2] = {
|
||||
static_cast<std::uint64_t>(right),
|
||||
static_cast<std::uint64_t>(right >> 64)};
|
||||
auto msg = simde_mm_load_si128(
|
||||
reinterpret_cast<const simde__m128i *>(lanes));
|
||||
msg = simde_mm_xor_si128(msg, seed);
|
||||
msg = simde_mm_xor_si128(msg,
|
||||
simde_mm_set_epi32(0, 0, 0, round + 1));
|
||||
const auto out = prg::aes128::eval(msg,
|
||||
static_cast<psnip_uint32_t>(round + 1));
|
||||
simde_mm_store_si128(reinterpret_cast<simde__m128i *>(lanes), out);
|
||||
wide f = lanes[0] | (wide{lanes[1]} << 64);
|
||||
f &= mask;
|
||||
left ^= f;
|
||||
const wide tmp = left;
|
||||
const mpf_word f = word_and(aes_prf(seed, right, round), mask);
|
||||
left.lo = left.lo ^ f.lo;
|
||||
left.hi = left.hi ^ f.hi;
|
||||
const mpf_word tmp = left;
|
||||
left = right;
|
||||
right = tmp;
|
||||
}
|
||||
val = (left << half) | right;
|
||||
val = word_or(word_shl(left, static_cast<unsigned>(half)), right);
|
||||
if (val < domain)
|
||||
return val;
|
||||
}
|
||||
throw prp_walk_error{};
|
||||
}
|
||||
|
||||
inline located locate(simde__m128i sigma, wide x, int hash,
|
||||
wide n, wide bucket_domain)
|
||||
inline located locate(simde__m128i sigma, mpf_word x, int hash,
|
||||
mpf_word n, mpf_word bucket_domain)
|
||||
{
|
||||
constexpr int kappa = 3;
|
||||
const wide y = permute(sigma,
|
||||
x + n * static_cast<unsigned>(hash), n * kappa);
|
||||
const mpf_word y = permute(sigma,
|
||||
word_add(x, word_mul_small(n, static_cast<std::uint64_t>(hash))),
|
||||
word_mul_small(n, kappa));
|
||||
const auto [quot, rem] = word_divmod(y, bucket_domain);
|
||||
if (quot.hi != uint256_t{0})
|
||||
throw std::runtime_error("multipoint bucket index does not fit");
|
||||
located out;
|
||||
out.bucket = static_cast<std::uint32_t>(y / bucket_domain);
|
||||
out.index = static_cast<std::uint64_t>(y % bucket_domain);
|
||||
out.bucket = static_cast<std::uint64_t>(quot.lo);
|
||||
out.index = rem;
|
||||
return out;
|
||||
}
|
||||
|
||||
inline std::uint32_t bucket_count_for(std::uint32_t t, std::uint32_t lambda)
|
||||
inline std::uint64_t bucket_count_for(std::uint64_t t, std::uint32_t lambda)
|
||||
{
|
||||
const double log2t = (t <= 1) ? 0.0 : std::log2(static_cast<double>(t));
|
||||
const double e = (static_cast<double>(lambda) + 130.0 + log2t) / 123.5;
|
||||
auto m = static_cast<std::uint32_t>(std::ceil(e * static_cast<double>(t)));
|
||||
auto m = static_cast<std::uint64_t>(std::ceil(e * static_cast<double>(t)));
|
||||
if (m < t + 1)
|
||||
m = t + 1;
|
||||
// Remark 1's simplification wants t ≥ 30. Below that, keep a 2t table.
|
||||
if (t < 30 && m < t * 2)
|
||||
m = t * 2;
|
||||
// At least κ buckets so each within-bucket index fits in the input type.
|
||||
if (m < 3)
|
||||
m = 3;
|
||||
return m;
|
||||
}
|
||||
|
||||
|
|
@ -199,28 +436,28 @@ inline std::uint32_t rng_seed(simde__m128i sigma)
|
|||
|
||||
struct slot
|
||||
{
|
||||
int item = -1;
|
||||
std::int64_t item = -1;
|
||||
int hash = -1;
|
||||
};
|
||||
|
||||
template <typename InputT>
|
||||
bool insert_cuckoo(simde__m128i sigma, const std::vector<InputT> & alphas,
|
||||
std::uint32_t m, wide n, wide bucket_domain,
|
||||
std::uint64_t m, mpf_word n, mpf_word bucket_domain,
|
||||
std::uint32_t max_evictions, std::vector<slot> & table)
|
||||
{
|
||||
table.assign(m, slot{});
|
||||
std::mt19937 rng(rng_seed(sigma));
|
||||
std::uniform_int_distribution<int> pick(0, 2);
|
||||
const int t = static_cast<int>(alphas.size());
|
||||
for (int omega = 0; omega < t; ++omega)
|
||||
const auto t = static_cast<std::int64_t>(alphas.size());
|
||||
for (std::int64_t omega = 0; omega < t; ++omega)
|
||||
{
|
||||
int cur = omega;
|
||||
std::int64_t cur = omega;
|
||||
int hash = pick(rng);
|
||||
std::uint32_t evictions = 0;
|
||||
for (;;)
|
||||
{
|
||||
const auto loc = locate(sigma,
|
||||
static_cast<wide>(alphas[static_cast<std::size_t>(cur)]),
|
||||
to_word(alphas[static_cast<std::size_t>(cur)]),
|
||||
hash, n, bucket_domain);
|
||||
if (loc.bucket >= m)
|
||||
return false;
|
||||
|
|
@ -229,7 +466,7 @@ bool insert_cuckoo(simde__m128i sigma, const std::vector<InputT> & alphas,
|
|||
table[loc.bucket] = slot{cur, hash};
|
||||
break;
|
||||
}
|
||||
const int evicted = table[loc.bucket].item;
|
||||
const std::int64_t evicted = table[loc.bucket].item;
|
||||
table[loc.bucket] = slot{cur, hash};
|
||||
cur = evicted;
|
||||
hash = pick(rng);
|
||||
|
|
@ -270,31 +507,26 @@ struct bucket_bare
|
|||
template <bool Verifiable,
|
||||
typename InteriorPRG,
|
||||
typename ExteriorPRG,
|
||||
typename BucketInput,
|
||||
typename InputT,
|
||||
typename OutputT>
|
||||
auto make_impl(std::vector<InputT> alphas, std::vector<OutputT> betas,
|
||||
multipoint_params params)
|
||||
{
|
||||
using bare = typename bucket_bare<Verifiable, InteriorPRG, ExteriorPRG,
|
||||
BucketInput, OutputT>::type;
|
||||
InputT, OutputT>::type;
|
||||
using key0 = multipoint_key<0, InputT, OutputT, bare>;
|
||||
using key1 = multipoint_key<1, InputT, OutputT, bare>;
|
||||
|
||||
static_assert(std::is_unsigned_v<InputT> && !std::is_same_v<InputT, bool>,
|
||||
static_assert(!std::is_same_v<InputT, bool>
|
||||
&& (std::is_unsigned_v<InputT> || std::is_same_v<InputT, uint256_t>),
|
||||
"make_multipoint: input domain must be an unsigned integer");
|
||||
static_assert(utils::bitlength_of_v<InputT> <= 32,
|
||||
"make_multipoint: input domain wider than 32 bits is not supported");
|
||||
static_assert(std::is_unsigned_v<BucketInput>
|
||||
&& !std::is_same_v<BucketInput, bool>,
|
||||
"make_multipoint: BucketInput must be an unsigned integer");
|
||||
static_assert(utils::bitlength_of_v<InputT> <= 256,
|
||||
"make_multipoint: input type is wider than a point key in this library");
|
||||
|
||||
if (alphas.size() != betas.size())
|
||||
throw std::invalid_argument("make_multipoint: point and payload counts differ");
|
||||
if (alphas.empty())
|
||||
throw std::invalid_argument("make_multipoint: no points");
|
||||
if (alphas.size() > static_cast<std::size_t>(std::numeric_limits<std::uint32_t>::max()))
|
||||
throw std::invalid_argument("make_multipoint: too many points");
|
||||
|
||||
{
|
||||
auto sorted = alphas;
|
||||
|
|
@ -303,19 +535,16 @@ auto make_impl(std::vector<InputT> alphas, std::vector<OutputT> betas,
|
|||
throw std::invalid_argument("make_multipoint: duplicate points");
|
||||
}
|
||||
|
||||
const auto t = static_cast<std::uint32_t>(alphas.size());
|
||||
const auto t = static_cast<std::uint64_t>(alphas.size());
|
||||
const auto m = bucket_count_for(t, params.lambda);
|
||||
constexpr std::size_t input_bits = utils::bitlength_of_v<InputT>;
|
||||
const wide n = domain_size(input_bits);
|
||||
const mpf_word n = domain_size(input_bits);
|
||||
constexpr int kappa = 3;
|
||||
const wide b = (n * kappa + m - 1) / m;
|
||||
constexpr std::size_t bucket_bits = utils::bitlength_of_v<BucketInput>;
|
||||
const wide bucket_cap = domain_size(bucket_bits);
|
||||
if (b > bucket_cap)
|
||||
{
|
||||
throw std::invalid_argument(
|
||||
"make_multipoint: bucket domain does not fit in BucketInput");
|
||||
}
|
||||
const mpf_word span = word_mul_small(n, kappa);
|
||||
const mpf_word den{uint256_t{m}, uint256_t{0}};
|
||||
const mpf_word numer = word_add(span,
|
||||
word_sub(den, mpf_word{uint256_t{1}, uint256_t{0}}));
|
||||
const mpf_word b = word_divmod(numer, den).first;
|
||||
|
||||
const int attempts = params.retries < 1 ? 1 : params.retries;
|
||||
for (int attempt = 0; attempt < attempts; ++attempt)
|
||||
|
|
@ -333,24 +562,26 @@ auto make_impl(std::vector<InputT> alphas, std::vector<OutputT> betas,
|
|||
right.sigma = sigma;
|
||||
left.bucket_count = m;
|
||||
right.bucket_count = m;
|
||||
left.bucket_domain = static_cast<std::uint64_t>(b);
|
||||
right.bucket_domain = static_cast<std::uint64_t>(b);
|
||||
left.buckets.reserve(m);
|
||||
right.buckets.reserve(m);
|
||||
left.bucket_domain = b;
|
||||
right.bucket_domain = b;
|
||||
left.buckets.reserve(static_cast<std::size_t>(m));
|
||||
right.buckets.reserve(static_cast<std::size_t>(m));
|
||||
|
||||
for (std::uint32_t i = 0; i < m; ++i)
|
||||
for (std::uint64_t i = 0; i < m; ++i)
|
||||
{
|
||||
BucketInput gamma{};
|
||||
InputT gamma{};
|
||||
OutputT beta{};
|
||||
if (table[i].item >= 0)
|
||||
if (table[static_cast<std::size_t>(i)].item >= 0)
|
||||
{
|
||||
const auto & alpha = alphas[static_cast<std::size_t>(table[i].item)];
|
||||
const auto loc = locate(sigma,
|
||||
static_cast<wide>(alpha), table[i].hash, n, b);
|
||||
const auto & alpha = alphas[static_cast<std::size_t>(
|
||||
table[static_cast<std::size_t>(i)].item)];
|
||||
const auto loc = locate(sigma, to_word(alpha),
|
||||
table[static_cast<std::size_t>(i)].hash, n, b);
|
||||
if (loc.bucket != i)
|
||||
throw prp_walk_error{};
|
||||
gamma = static_cast<BucketInput>(loc.index);
|
||||
beta = betas[static_cast<std::size_t>(table[i].item)];
|
||||
gamma = from_word<InputT>(loc.index);
|
||||
beta = betas[static_cast<std::size_t>(
|
||||
table[static_cast<std::size_t>(i)].item)];
|
||||
}
|
||||
auto made = make_bucket<Verifiable, InteriorPRG, ExteriorPRG>(
|
||||
gamma, beta);
|
||||
|
|
@ -371,6 +602,7 @@ inline void absorb_proof(proof_token & acc, const proof_token & inner)
|
|||
{
|
||||
acc = detail::vdpf::xor_proof(acc, inner);
|
||||
acc[0] = detail::vdpf::mmo(acc[0], 1);
|
||||
acc[1] = detail::vdpf::mmo(acc[1], 2);
|
||||
}
|
||||
|
||||
template <typename Key>
|
||||
|
|
@ -380,17 +612,17 @@ typename Key::share_type eval_at(const Key & key, typename Key::input_type x,
|
|||
using input_type = typename Key::input_type;
|
||||
using bucket_input = typename Key::bucket_input;
|
||||
constexpr std::size_t input_bits = utils::bitlength_of_v<input_type>;
|
||||
const wide n = domain_size(input_bits);
|
||||
const wide b = key.bucket_domain;
|
||||
const mpf_word n = domain_size(input_bits);
|
||||
const mpf_word b = key.bucket_domain;
|
||||
typename Key::share_type sum =
|
||||
Key::share_type::from_raw(typename Key::output_type{});
|
||||
|
||||
for (int hash = 0; hash < static_cast<int>(Key::kappa); ++hash)
|
||||
{
|
||||
const auto loc = locate(key.sigma, static_cast<wide>(x), hash, n, b);
|
||||
const auto loc = locate(key.sigma, to_word(x), hash, n, b);
|
||||
if (loc.bucket >= key.bucket_count)
|
||||
throw std::runtime_error("multipoint eval: bucket out of range");
|
||||
const auto gamma = static_cast<bucket_input>(loc.index);
|
||||
const auto gamma = from_word<bucket_input>(loc.index);
|
||||
const auto & bucket = key.buckets[loc.bucket];
|
||||
if constexpr (Key::is_verifiable)
|
||||
{
|
||||
|
|
@ -413,7 +645,6 @@ typename Key::share_type eval_at(const Key & key, typename Key::input_type x,
|
|||
/// @brief Cuckoo-pack distinct points into ordinary point-key buckets.
|
||||
/// @tparam InteriorPRG PRG that expands interior nodes. Defaults to `dpf::prg::aes128`
|
||||
/// @tparam ExteriorPRG PRG that expands the root. Defaults to `InteriorPRG`
|
||||
/// @tparam BucketInput unsigned type of a bucket index. Defaults to `uint32_t`
|
||||
/// @tparam AlphaRange range of distinct domain points
|
||||
/// @tparam BetaRange range of payloads, one per point
|
||||
/// @param alphas the secret points
|
||||
|
|
@ -421,11 +652,12 @@ typename Key::share_type eval_at(const Key & key, typename Key::input_type x,
|
|||
/// @param params packing knobs. `lambda` is the Remark 1 failure target
|
||||
/// @return the two party keys
|
||||
/// @throws std::invalid_argument if the lists differ in length, are empty,
|
||||
/// contain a duplicate, or a bucket index does not fit `BucketInput`
|
||||
/// or contain a duplicate
|
||||
/// @throws std::runtime_error if cuckoo hashing does not succeed
|
||||
/// @note Following de Castro and Polychroniadou, EUROCRYPT 2022, §4 (ePrint 2021/580): κ = 3 cuckoo buckets, one point key each.
|
||||
/// \complexity For t points, `bucket_count_for` sets m = ceil((lambda + 130 + log2(t)) / 123.5 * t) (at least 2t when t < 30, and at least 3). Each attempt inserts t cuckoo items (up to `max_evictions` swaps each) and then one `make_dpf` per bucket. The within-bucket domain `b` is ceil(3 * 2^{input bits} / m). Counted `insert_cuckoo` and the bucket loop. Retries are `params.retries`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename BucketInput = std::uint32_t,
|
||||
typename AlphaRange,
|
||||
typename BetaRange>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
|
|
@ -434,7 +666,7 @@ auto make_multipoint(const AlphaRange & alphas, const BetaRange & betas,
|
|||
{
|
||||
using input_type = std::decay_t<decltype(*std::begin(alphas))>;
|
||||
using output_type = std::decay_t<decltype(*std::begin(betas))>;
|
||||
return detail::mpf::make_impl<false, InteriorPRG, ExteriorPRG, BucketInput>(
|
||||
return detail::mpf::make_impl<false, InteriorPRG, ExteriorPRG>(
|
||||
std::vector<input_type>(std::begin(alphas), std::end(alphas)),
|
||||
std::vector<output_type>(std::begin(betas), std::end(betas)),
|
||||
params);
|
||||
|
|
@ -447,11 +679,12 @@ auto make_multipoint(const AlphaRange & alphas, const BetaRange & betas,
|
|||
/// @param params packing knobs
|
||||
/// @return the two verifiable party keys
|
||||
/// @throws std::invalid_argument if the lists differ in length, are empty,
|
||||
/// contain a duplicate, or a bucket index does not fit `BucketInput`
|
||||
/// or contain a duplicate
|
||||
/// @throws std::runtime_error if cuckoo hashing does not succeed
|
||||
/// @note Following de Castro and Polychroniadou, EUROCRYPT 2022, §4 (ePrint 2021/580): κ = 3 cuckoo buckets, one point key each.
|
||||
/// \complexity For t points, `bucket_count_for` sets m = ceil((lambda + 130 + log2(t)) / 123.5 * t) (at least 2t when t < 30, and at least 3). Each attempt inserts t cuckoo items (up to `max_evictions` swaps each) and then one `make_dpf` per bucket. The within-bucket domain `b` is ceil(3 * 2^{input bits} / m). Counted `insert_cuckoo` and the bucket loop. Retries are `params.retries`.
|
||||
template <typename InteriorPRG = dpf::prg::aes128,
|
||||
typename ExteriorPRG = InteriorPRG,
|
||||
typename BucketInput = std::uint32_t,
|
||||
typename AlphaRange,
|
||||
typename BetaRange>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
|
|
@ -460,7 +693,7 @@ auto make_multipoint(const AlphaRange & alphas, const BetaRange & betas,
|
|||
{
|
||||
using input_type = std::decay_t<decltype(*std::begin(alphas))>;
|
||||
using output_type = std::decay_t<decltype(*std::begin(betas))>;
|
||||
return detail::mpf::make_impl<true, InteriorPRG, ExteriorPRG, BucketInput>(
|
||||
return detail::mpf::make_impl<true, InteriorPRG, ExteriorPRG>(
|
||||
std::vector<input_type>(std::begin(alphas), std::end(alphas)),
|
||||
std::vector<output_type>(std::begin(betas), std::end(betas)),
|
||||
params);
|
||||
|
|
@ -472,6 +705,7 @@ auto make_multipoint(const AlphaRange & alphas, const BetaRange & betas,
|
|||
/// @param x the query point
|
||||
/// @return the party's share of the payload, or of zero off the packed points
|
||||
/// @throws std::runtime_error if a located bucket is outside the key
|
||||
/// \complexity Three `eval_point` calls (`Key::kappa` is 3), each O(n_b) where n_b is the bucket key depth.
|
||||
template <typename Key,
|
||||
std::enable_if_t<is_multipoint_key_v<Key>, int> = 0>
|
||||
auto eval_multipoint(const Key & key, typename Key::input_type x)
|
||||
|
|
@ -486,6 +720,7 @@ auto eval_multipoint(const Key & key, typename Key::input_type x)
|
|||
/// @param pr proof token replaced with this query's folded proof
|
||||
/// @return the party's share of the payload
|
||||
/// @throws std::runtime_error if a located bucket is outside the key
|
||||
/// \complexity Three `eval_point` calls (`Key::kappa` is 3), each O(n_b) where n_b is the bucket key depth.
|
||||
template <typename Key,
|
||||
std::enable_if_t<is_multipoint_key_v<Key>, int> = 0>
|
||||
auto eval_multipoint(const Key & key, typename Key::input_type x, prove_ref pr)
|
||||
|
|
@ -504,6 +739,7 @@ auto eval_multipoint(const Key & key, typename Key::input_type x, prove_ref pr)
|
|||
/// @param xs the query points
|
||||
/// @param out where each share is written
|
||||
/// @throws std::runtime_error if a located bucket is outside the key
|
||||
/// \complexity Three `eval_point` calls (`Key::kappa` is 3), each O(n_b) where n_b is the bucket key depth.
|
||||
template <typename Key, typename Range, typename OutIt,
|
||||
std::enable_if_t<is_multipoint_key_v<Key>, int> = 0>
|
||||
void eval_multipoint(const Key & key, const Range & xs, OutIt out)
|
||||
|
|
@ -521,6 +757,7 @@ void eval_multipoint(const Key & key, const Range & xs, OutIt out)
|
|||
/// @param out where each share is written
|
||||
/// @param pr proof token replaced with the folded proof of `xs`
|
||||
/// @throws std::runtime_error if a located bucket is outside the key
|
||||
/// \complexity Three `eval_point` calls (`Key::kappa` is 3), each O(n_b) where n_b is the bucket key depth.
|
||||
template <typename Key, typename Range, typename OutIt,
|
||||
std::enable_if_t<is_multipoint_key_v<Key>, int> = 0>
|
||||
void eval_multipoint(const Key & key, const Range & xs, OutIt out, prove_ref pr)
|
||||
|
|
|
|||
43
include/dpf/net/asio_ns.hpp
Normal file
43
include/dpf/net/asio_ns.hpp
Normal file
|
|
@ -0,0 +1,43 @@
|
|||
/// @file dpf/net/asio_ns.hpp
|
||||
/// @brief Standalone asio as `asio::` inside the networking namespaces.
|
||||
/// @details `dpf/asio.hpp` declares `dpf::asio` (DPF key shipping), which
|
||||
/// would otherwise capture unqualified `asio::` lookups made inside
|
||||
/// `namespace dpf` once both are included, in either order.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_NET_ASIO_NS_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_NET_ASIO_NS_HPP__
|
||||
|
||||
#include <asio.hpp>
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace net
|
||||
{
|
||||
namespace asio = ::asio;
|
||||
} // namespace net
|
||||
namespace protocol
|
||||
{
|
||||
namespace asio = ::asio;
|
||||
} // namespace protocol
|
||||
namespace session
|
||||
{
|
||||
namespace asio = ::asio;
|
||||
} // namespace session
|
||||
namespace run
|
||||
{
|
||||
namespace asio = ::asio;
|
||||
} // namespace run
|
||||
namespace app
|
||||
{
|
||||
namespace asio = ::asio;
|
||||
} // namespace app
|
||||
namespace async
|
||||
{
|
||||
namespace asio = ::asio;
|
||||
} // namespace async
|
||||
namespace factory
|
||||
{
|
||||
namespace asio = ::asio;
|
||||
} // namespace factory
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_NET_ASIO_NS_HPP__
|
||||
876
include/dpf/net/async_round_sink.hpp
Normal file
876
include/dpf/net/async_round_sink.hpp
Normal file
|
|
@ -0,0 +1,876 @@
|
|||
/// @file dpf/net/async_round_sink.hpp
|
||||
/// @brief Event-driven `RoundSink` over an `async_stream_array`.
|
||||
/// @details Both ends first exchange a hello on lane 0 carrying the plan shape
|
||||
/// (rounds, slot widths, lanes, instances, framing) and an epoch. A
|
||||
/// mismatch fails both ends with a message that names the field.
|
||||
/// Unframed lanes carry one round each (round == lane). Framed lanes
|
||||
/// carry `{round, nbytes}` headers; payloads are read straight into the
|
||||
/// round's inbox, and a partial prefix of instances is ready as soon
|
||||
/// as it lands. Every outbound round stays in memory, so after a
|
||||
/// transport error the sink can take a replacement link from
|
||||
/// `sink_options::reconnect`, exchange what each side received, and
|
||||
/// resend only the missing bytes. The schedule above never rewinds.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_NET_ASYNC_ROUND_SINK_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_NET_ASYNC_ROUND_SINK_HPP__
|
||||
|
||||
#include <algorithm>
|
||||
#include <array>
|
||||
#include <atomic>
|
||||
#include <chrono>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <functional>
|
||||
#include <memory>
|
||||
#include <mutex>
|
||||
#include <stdexcept>
|
||||
#include <string>
|
||||
#include <system_error>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/net/asio_ns.hpp"
|
||||
|
||||
#include "dpf/net/async_stream_array.hpp"
|
||||
#include "dpf/net/policy.hpp"
|
||||
#include "dpf/net/round_lane.hpp"
|
||||
#include "dpf/net/round_sink.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace net
|
||||
{
|
||||
|
||||
/// @brief Construction options for `async_round_sink`.
|
||||
struct sink_options
|
||||
{
|
||||
framing_mode framing = framing_mode::automatic;
|
||||
/// Exchange the plan-shape hello. Both ends must agree.
|
||||
bool hello = true;
|
||||
/// Bound on waiting for a full write window to drain.
|
||||
std::chrono::milliseconds drain_timeout{30000};
|
||||
/// Bound on the hello exchange after a reconnect.
|
||||
std::chrono::milliseconds handshake_timeout{30000};
|
||||
/// On a transport error, return a replacement link (the peer must replace
|
||||
/// its end too) or nullptr to fail. Runs on the drive thread.
|
||||
std::function<async_stream_array *(const std::error_code &)> reconnect;
|
||||
unsigned max_reconnects = 3;
|
||||
};
|
||||
|
||||
/// @brief Per-sink counters.
|
||||
struct sink_stats
|
||||
{
|
||||
std::uint64_t bytes_sent = 0;
|
||||
std::uint64_t bytes_received = 0;
|
||||
std::uint64_t flushes = 0;
|
||||
std::uint64_t window_waits = 0;
|
||||
std::uint64_t window_wait_ns = 0;
|
||||
std::uint64_t resumes = 0;
|
||||
std::uint64_t resent_bytes = 0;
|
||||
bool framed = false;
|
||||
std::size_t lanes = 0;
|
||||
};
|
||||
|
||||
/// @brief `RoundSink` whose I/O completes through an `io_context`, never spins.
|
||||
class async_round_sink : public RoundSink
|
||||
{
|
||||
public:
|
||||
async_round_sink(async_stream_array & streams,
|
||||
std::vector<std::size_t> slot_bytes_per_round, std::size_t count = 1,
|
||||
sink_options opt = {})
|
||||
: core_(std::make_shared<core>(streams, streams,
|
||||
std::move(slot_bytes_per_round), count, std::move(opt)))
|
||||
{
|
||||
core_->start();
|
||||
}
|
||||
|
||||
/// @brief Send on `out`, receive on `in` (a ring edge: to the previous
|
||||
/// party, from the next). Both must run on the same `io_context`.
|
||||
/// Reconnect resume needs a single link.
|
||||
async_round_sink(async_stream_array & out, async_stream_array & in,
|
||||
std::vector<std::size_t> slot_bytes_per_round, std::size_t count = 1,
|
||||
sink_options opt = {})
|
||||
: core_(std::make_shared<core>(out, in, std::move(slot_bytes_per_round),
|
||||
count, std::move(opt)))
|
||||
{
|
||||
core_->start();
|
||||
}
|
||||
|
||||
async_round_sink(const async_round_sink &) = delete;
|
||||
async_round_sink & operator=(const async_round_sink &) = delete;
|
||||
|
||||
~async_round_sink() override
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(core_->mu);
|
||||
core_->dead = true;
|
||||
}
|
||||
|
||||
std::size_t count() const noexcept override { return core_->count; }
|
||||
std::size_t rounds() const noexcept override { return core_->slots.size(); }
|
||||
bool framed() const noexcept { return core_->map.framed; }
|
||||
std::size_t lanes() const noexcept { return core_->map.n_lanes; }
|
||||
|
||||
std::size_t slot_bytes(std::uint16_t round) const override
|
||||
{
|
||||
if (round >= core_->slots.size())
|
||||
throw std::out_of_range("async_round_sink round");
|
||||
return core_->slots[round];
|
||||
}
|
||||
|
||||
asio::io_context & context() noexcept { return *core_->io; }
|
||||
async_stream_array & link() noexcept { return *core_->streams; }
|
||||
async_stream_array & in_link() noexcept { return *core_->in; }
|
||||
|
||||
sink_stats stats() const
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(core_->mu);
|
||||
sink_stats s = core_->st;
|
||||
s.framed = core_->map.framed;
|
||||
s.lanes = core_->map.n_lanes;
|
||||
return s;
|
||||
}
|
||||
|
||||
void submit(std::uint16_t round, std::size_t index,
|
||||
const std::uint8_t * bytes, std::size_t n) override
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(core_->mu);
|
||||
core_->window(round).submit(index, bytes, n);
|
||||
}
|
||||
|
||||
bool peer_ready(std::uint16_t round, std::size_t index) const override
|
||||
{
|
||||
{
|
||||
// Bytes that already arrived stay readable after a link error.
|
||||
std::lock_guard<std::mutex> lock(core_->mu);
|
||||
if (core_->ready_locked(round, index))
|
||||
return true;
|
||||
}
|
||||
core_->raise_if_failed();
|
||||
std::lock_guard<std::mutex> lock(core_->mu);
|
||||
return core_->ready_locked(round, index);
|
||||
}
|
||||
|
||||
void read_peer(std::uint16_t round, std::size_t index, std::uint8_t * out,
|
||||
std::size_t n) const override
|
||||
{
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(core_->mu);
|
||||
if (round < core_->slots.size() && core_->ready_locked(round, index))
|
||||
{
|
||||
if (n != core_->slots[round])
|
||||
throw std::invalid_argument("async_round_sink read size");
|
||||
if (n != 0)
|
||||
std::memcpy(out, core_->inbox[round].data() + index * n, n);
|
||||
return;
|
||||
}
|
||||
}
|
||||
core_->raise_if_failed();
|
||||
std::lock_guard<std::mutex> lock(core_->mu);
|
||||
if (round >= core_->slots.size())
|
||||
throw std::out_of_range("async_round_sink read_peer round");
|
||||
if (!core_->ready_locked(round, index))
|
||||
throw std::logic_error("async_round_sink: peer not ready");
|
||||
if (n != core_->slots[round])
|
||||
throw std::invalid_argument("async_round_sink read size");
|
||||
if (n != 0)
|
||||
std::memcpy(out, core_->inbox[round].data() + index * n, n);
|
||||
}
|
||||
|
||||
void flush() override
|
||||
{
|
||||
for (std::uint16_t r = 0; r < core_->slots.size(); ++r)
|
||||
{
|
||||
bool due = false;
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(core_->mu);
|
||||
const auto & w = core_->win[r];
|
||||
due = w.next_unwritten() > w.flushed()
|
||||
|| (core_->slots[r] == 0 && !core_->announced[r]);
|
||||
}
|
||||
if (due)
|
||||
flush_round(r);
|
||||
}
|
||||
}
|
||||
|
||||
void flush_round(std::uint16_t round) override
|
||||
{
|
||||
core_->raise_if_failed();
|
||||
const std::size_t lane = core_->flush_one(round);
|
||||
core_->await_window(lane);
|
||||
}
|
||||
|
||||
void poll() override
|
||||
{
|
||||
core_->restart_if_stopped();
|
||||
core_->io->poll();
|
||||
}
|
||||
|
||||
bool wait_io() override
|
||||
{
|
||||
core_->raise_if_failed();
|
||||
core_->restart_if_stopped();
|
||||
const bool ran = core_->io->run_one() > 0;
|
||||
core_->raise_if_failed();
|
||||
return ran;
|
||||
}
|
||||
|
||||
bool wait_io_for(std::chrono::milliseconds budget) override
|
||||
{
|
||||
core_->raise_if_failed();
|
||||
core_->restart_if_stopped();
|
||||
const bool ran = budget.count() <= 0 ? core_->io->poll_one() > 0
|
||||
: core_->io->run_one_for(budget) > 0;
|
||||
core_->raise_if_failed();
|
||||
return ran;
|
||||
}
|
||||
|
||||
bool can_block() const noexcept override { return true; }
|
||||
|
||||
bool can_send_ahead() const noexcept override
|
||||
{
|
||||
for (std::size_t l = 0; l < core_->map.n_lanes; ++l)
|
||||
{
|
||||
const std::size_t w = core_->streams->lane_window_bytes(l);
|
||||
if (w != 0 && core_->streams->lane_buffered_bytes(l) > w)
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
std::uint64_t progress() const noexcept override
|
||||
{
|
||||
return core_->progress.load(std::memory_order_relaxed);
|
||||
}
|
||||
|
||||
const void * wait_domain() const noexcept override { return core_->io; }
|
||||
|
||||
private:
|
||||
static constexpr std::uint32_t k_magic = 0x48535044u; // 'DPSH'
|
||||
static constexpr std::uint16_t k_version = 1;
|
||||
static constexpr std::size_t k_hello_fixed = 36;
|
||||
static constexpr std::uint32_t k_zero_seen = 1;
|
||||
|
||||
struct core : std::enable_shared_from_this<core>
|
||||
{
|
||||
core(async_stream_array & s, async_stream_array & rx,
|
||||
std::vector<std::size_t> slots_, std::size_t count_, sink_options opt_)
|
||||
: streams(&s),
|
||||
in(&rx),
|
||||
io(&s.context()),
|
||||
count(count_),
|
||||
slots(std::move(slots_)),
|
||||
map(std::min(s.size(), rx.size()), slots.size(), opt_.framing),
|
||||
opt(std::move(opt_))
|
||||
{
|
||||
if (count == 0)
|
||||
throw std::invalid_argument("async_round_sink: count 0");
|
||||
if (s.size() == 0 || rx.size() == 0)
|
||||
throw std::invalid_argument("async_round_sink: empty streams");
|
||||
if (&s.context() != &rx.context())
|
||||
throw std::invalid_argument(
|
||||
"async_round_sink: out and in links need one io_context");
|
||||
if (&s != &rx && opt.reconnect)
|
||||
throw std::invalid_argument(
|
||||
"async_round_sink: reconnect needs a single link");
|
||||
const std::size_t nr = slots.size();
|
||||
win.reserve(nr);
|
||||
inbox.resize(nr);
|
||||
in_filled.assign(nr, 0);
|
||||
zero_seen.assign(nr, false);
|
||||
announced.assign(nr, false);
|
||||
legacy_started.assign(nr, false);
|
||||
legacy_wanted.assign(nr, false);
|
||||
for (std::size_t k = 0; k < nr; ++k)
|
||||
{
|
||||
win.emplace_back(count, slots[k]);
|
||||
inbox[k].assign(count * slots[k], 0);
|
||||
}
|
||||
fingerprint = slots_fingerprint(slots);
|
||||
}
|
||||
|
||||
// --- lifecycle -------------------------------------------------
|
||||
void start()
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(mu);
|
||||
if (opt.hello)
|
||||
begin_hello_locked();
|
||||
else
|
||||
hello_done_locked();
|
||||
}
|
||||
|
||||
void restart_if_stopped()
|
||||
{
|
||||
if (io->stopped())
|
||||
io->restart();
|
||||
}
|
||||
|
||||
// --- hello ------------------------------------------------------
|
||||
std::shared_ptr<std::vector<std::uint8_t>> build_hello_locked() const
|
||||
{
|
||||
const std::size_t r = slots.size();
|
||||
auto buf = acquire_buffer(k_hello_fixed + 4 * r);
|
||||
auto * p = buf->data();
|
||||
std::memset(p, 0, buf->size());
|
||||
detail::put_u32(p + 0, k_magic);
|
||||
p[4] = static_cast<std::uint8_t>(k_version & 0xffu);
|
||||
p[5] = static_cast<std::uint8_t>(k_version >> 8);
|
||||
p[6] = map.framed ? 1 : 0;
|
||||
detail::put_u32(p + 8, static_cast<std::uint32_t>(r));
|
||||
detail::put_u32(p + 12, static_cast<std::uint32_t>(map.n_lanes));
|
||||
detail::put_u32(p + 16, static_cast<std::uint32_t>(count));
|
||||
detail::put_u32(p + 20, gen);
|
||||
for (int i = 0; i < 8; ++i)
|
||||
p[24 + i] = static_cast<std::uint8_t>((fingerprint >> (8 * i)) & 0xffu);
|
||||
for (std::size_t k = 0; k < r; ++k)
|
||||
{
|
||||
const std::uint32_t got = slots[k] == 0
|
||||
? (zero_seen[k] ? k_zero_seen : 0)
|
||||
: static_cast<std::uint32_t>(in_filled[k]);
|
||||
detail::put_u32(p + k_hello_fixed + 4 * k, got);
|
||||
}
|
||||
return buf;
|
||||
}
|
||||
|
||||
void begin_hello_locked()
|
||||
{
|
||||
hello_ok = false;
|
||||
auto self = shared_from_this();
|
||||
const std::uint32_t g = gen;
|
||||
streams->async_write_owned(0, build_hello_locked(),
|
||||
[self, g](const std::error_code & ec) { self->on_write(g, ec); });
|
||||
hello_in = std::make_shared<std::vector<std::uint8_t>>(
|
||||
k_hello_fixed + 4 * slots.size());
|
||||
auto buf = hello_in;
|
||||
in->async_read(0, buf->data(), buf->size(),
|
||||
[self, g, buf](const std::error_code & ec) {
|
||||
self->on_hello(g, buf, ec);
|
||||
});
|
||||
}
|
||||
|
||||
void on_hello(std::uint32_t g,
|
||||
const std::shared_ptr<std::vector<std::uint8_t>> & buf,
|
||||
const std::error_code & ec)
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(mu);
|
||||
if (dead || g != gen)
|
||||
return;
|
||||
if (ec)
|
||||
{
|
||||
fail_locked(ec, "hello read", false);
|
||||
return;
|
||||
}
|
||||
const auto * p = buf->data();
|
||||
if (detail::get_u32(p) != k_magic)
|
||||
{
|
||||
fail_locked(std::make_error_code(std::errc::protocol_error),
|
||||
"peer did not send a sink hello (is it an async_round_sink "
|
||||
"or stream_array_sink with hello enabled?)", true);
|
||||
return;
|
||||
}
|
||||
const std::uint16_t ver = static_cast<std::uint16_t>(p[4] | (p[5] << 8));
|
||||
const bool peer_framed = (p[6] & 1) != 0;
|
||||
const std::uint32_t pr = detail::get_u32(p + 8);
|
||||
const std::uint32_t pl = detail::get_u32(p + 12);
|
||||
const std::uint32_t pc = detail::get_u32(p + 16);
|
||||
const std::uint32_t pg = detail::get_u32(p + 20);
|
||||
std::uint64_t pf = 0;
|
||||
for (int i = 0; i < 8; ++i)
|
||||
pf |= static_cast<std::uint64_t>(p[24 + i]) << (8 * i);
|
||||
std::string why;
|
||||
if (ver != k_version)
|
||||
why += " version " + std::to_string(ver) + " vs "
|
||||
+ std::to_string(k_version) + ";";
|
||||
if (pr != slots.size())
|
||||
why += " rounds " + std::to_string(pr) + " vs "
|
||||
+ std::to_string(slots.size()) + ";";
|
||||
else if (pf != fingerprint)
|
||||
why += " slot widths differ;";
|
||||
if (pl != map.n_lanes)
|
||||
why += " lanes " + std::to_string(pl) + " vs "
|
||||
+ std::to_string(map.n_lanes) + ";";
|
||||
if (pc != count)
|
||||
why += " instances " + std::to_string(pc) + " vs "
|
||||
+ std::to_string(count) + ";";
|
||||
if (peer_framed != map.framed)
|
||||
why += std::string(" framing ") + (peer_framed ? "on" : "off")
|
||||
+ " vs " + (map.framed ? "on" : "off") + ";";
|
||||
if (pg != gen)
|
||||
why += " epoch " + std::to_string(pg) + " vs "
|
||||
+ std::to_string(gen) + ";";
|
||||
if (!why.empty())
|
||||
{
|
||||
fail_locked(std::make_error_code(std::errc::protocol_error),
|
||||
"peer sink disagrees (peer vs this side):" + why, true);
|
||||
return;
|
||||
}
|
||||
std::vector<std::uint32_t> peer_got(slots.size());
|
||||
for (std::size_t k = 0; k < slots.size(); ++k)
|
||||
peer_got[k] = detail::get_u32(p + k_hello_fixed + 4 * k);
|
||||
if (gen != 0 && !resend_locked(peer_got))
|
||||
return;
|
||||
hello_done_locked();
|
||||
}
|
||||
|
||||
void hello_done_locked()
|
||||
{
|
||||
hello_ok = true;
|
||||
if (map.framed)
|
||||
{
|
||||
hdr_bufs.clear();
|
||||
for (std::size_t lane = 0; lane < map.n_lanes; ++lane)
|
||||
{
|
||||
hdr_bufs.push_back(
|
||||
std::make_shared<std::array<std::uint8_t, round_lane_hdr::size>>());
|
||||
read_lane_hdr_locked(lane);
|
||||
}
|
||||
}
|
||||
else
|
||||
{
|
||||
for (std::size_t r = 0; r < slots.size(); ++r)
|
||||
if (legacy_wanted[r])
|
||||
issue_legacy_locked(static_cast<std::uint16_t>(r));
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief After a reconnect: send what the peer reported missing.
|
||||
bool resend_locked(const std::vector<std::uint32_t> & peer_got)
|
||||
{
|
||||
auto self = shared_from_this();
|
||||
const std::uint32_t g = gen;
|
||||
for (std::size_t k = 0; k < slots.size(); ++k)
|
||||
{
|
||||
const std::size_t sb = slots[k];
|
||||
const auto round = static_cast<std::uint16_t>(k);
|
||||
const std::size_t lane = map.lane(round);
|
||||
if (sb == 0)
|
||||
{
|
||||
if (map.framed && announced[k] && peer_got[k] != k_zero_seen)
|
||||
{
|
||||
auto frame = acquire_buffer(round_lane_hdr::size);
|
||||
round_lane_hdr{round, 0}.pack(frame->data());
|
||||
streams->async_write_owned(lane, std::move(frame),
|
||||
[self, g](const std::error_code & ec) {
|
||||
self->on_write(g, ec);
|
||||
});
|
||||
}
|
||||
continue;
|
||||
}
|
||||
const std::size_t sent = win[k].flushed() * sb;
|
||||
const std::size_t got = peer_got[k];
|
||||
if (got > sent || got % sb != 0)
|
||||
{
|
||||
fail_locked(std::make_error_code(std::errc::protocol_error),
|
||||
"resume: peer reports " + std::to_string(got)
|
||||
+ " bytes of round " + std::to_string(k)
|
||||
+ ", this side sent " + std::to_string(sent),
|
||||
true);
|
||||
return false;
|
||||
}
|
||||
if (got == sent)
|
||||
continue;
|
||||
const std::size_t nbytes = sent - got;
|
||||
const std::uint8_t * src = win[k].out_at(got / sb);
|
||||
st.resent_bytes += nbytes;
|
||||
if (map.framed)
|
||||
{
|
||||
auto frame = acquire_buffer(round_lane_hdr::size + nbytes);
|
||||
round_lane_hdr{round, static_cast<std::uint32_t>(nbytes)}.pack(
|
||||
frame->data());
|
||||
std::memcpy(frame->data() + round_lane_hdr::size, src, nbytes);
|
||||
streams->async_write_owned(lane, std::move(frame),
|
||||
[self, g](const std::error_code & ec) {
|
||||
self->on_write(g, ec);
|
||||
});
|
||||
}
|
||||
else
|
||||
{
|
||||
streams->async_write(lane, src, nbytes,
|
||||
[self, g](const std::error_code & ec) {
|
||||
self->on_write(g, ec);
|
||||
});
|
||||
}
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
// --- errors and resume ------------------------------------------
|
||||
void fail_locked(const std::error_code & ec, const std::string & what,
|
||||
bool is_fatal)
|
||||
{
|
||||
if (fail)
|
||||
return;
|
||||
fail = ec;
|
||||
fail_what = "async_round_sink: " + what;
|
||||
fatal = is_fatal;
|
||||
}
|
||||
|
||||
void on_write(std::uint32_t g, const std::error_code & ec)
|
||||
{
|
||||
if (!ec)
|
||||
return;
|
||||
std::lock_guard<std::mutex> lock(mu);
|
||||
if (dead || g != gen)
|
||||
return;
|
||||
fail_locked(ec, "write", false);
|
||||
}
|
||||
|
||||
void raise_if_failed()
|
||||
{
|
||||
std::error_code ec;
|
||||
std::string what;
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(mu);
|
||||
if (!fail)
|
||||
return;
|
||||
if (fatal || !opt.reconnect || reconnects >= opt.max_reconnects)
|
||||
throw std::system_error(fail, fail_what);
|
||||
ec = fail;
|
||||
what = fail_what;
|
||||
}
|
||||
async_stream_array * next = nullptr;
|
||||
try
|
||||
{
|
||||
next = opt.reconnect(ec);
|
||||
}
|
||||
catch (const std::exception & e)
|
||||
{
|
||||
throw std::system_error(ec, what + "; reconnect failed: " + e.what());
|
||||
}
|
||||
if (next == nullptr)
|
||||
throw std::system_error(ec, what + "; no replacement link");
|
||||
resume(*next);
|
||||
}
|
||||
|
||||
void resume(async_stream_array & next)
|
||||
{
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(mu);
|
||||
if (next.size() < map.n_lanes)
|
||||
throw std::invalid_argument(
|
||||
"async_round_sink resume: replacement link has fewer lanes");
|
||||
++gen;
|
||||
++reconnects;
|
||||
++st.resumes;
|
||||
fail.clear();
|
||||
fail_what.clear();
|
||||
fatal = false;
|
||||
streams = &next;
|
||||
in = &next;
|
||||
io = &next.context();
|
||||
for (std::size_t r = 0; r < slots.size(); ++r)
|
||||
{
|
||||
if (legacy_started[r] && in_filled[r] < count * slots[r])
|
||||
{
|
||||
legacy_started[r] = false;
|
||||
legacy_wanted[r] = true;
|
||||
}
|
||||
}
|
||||
begin_hello_locked();
|
||||
}
|
||||
const auto deadline = std::chrono::steady_clock::now()
|
||||
+ opt.handshake_timeout;
|
||||
for (;;)
|
||||
{
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(mu);
|
||||
if (hello_ok)
|
||||
return;
|
||||
if (fail)
|
||||
throw std::system_error(fail, fail_what + " (during resume)");
|
||||
}
|
||||
const auto now = std::chrono::steady_clock::now();
|
||||
if (now >= deadline)
|
||||
throw std::system_error(std::make_error_code(std::errc::timed_out),
|
||||
"async_round_sink: resume hello timed out");
|
||||
restart_if_stopped();
|
||||
io->run_one_for(std::min<std::chrono::milliseconds>(
|
||||
std::chrono::duration_cast<std::chrono::milliseconds>(
|
||||
deadline - now),
|
||||
std::chrono::milliseconds(10)));
|
||||
}
|
||||
}
|
||||
|
||||
// --- reads ------------------------------------------------------
|
||||
/// @brief Every round's peer bytes (and zero-width announcements) are in.
|
||||
bool complete_locked() const
|
||||
{
|
||||
for (std::size_t r = 0; r < slots.size(); ++r)
|
||||
{
|
||||
const std::size_t sb = slots[r];
|
||||
if (sb == 0 ? (map.framed && !zero_seen[r])
|
||||
: in_filled[r] < count * sb)
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool ready_locked(std::uint16_t round, std::size_t index)
|
||||
{
|
||||
if (round >= slots.size())
|
||||
return false;
|
||||
const std::size_t sb = slots[round];
|
||||
if (map.framed)
|
||||
{
|
||||
if (sb == 0)
|
||||
return zero_seen[round];
|
||||
return in_filled[round] >= (index + 1) * sb;
|
||||
}
|
||||
if (sb == 0)
|
||||
return true;
|
||||
if (!legacy_started[round])
|
||||
{
|
||||
if (hello_ok)
|
||||
issue_legacy_locked(round);
|
||||
else
|
||||
legacy_wanted[round] = true;
|
||||
}
|
||||
return in_filled[round] == count * sb;
|
||||
}
|
||||
|
||||
void issue_legacy_locked(std::uint16_t round)
|
||||
{
|
||||
if (legacy_started[round])
|
||||
return;
|
||||
legacy_started[round] = true;
|
||||
legacy_wanted[round] = false;
|
||||
const std::size_t bytes = count * slots[round];
|
||||
if (bytes == 0)
|
||||
{
|
||||
in_filled[round] = 0;
|
||||
return;
|
||||
}
|
||||
auto self = shared_from_this();
|
||||
const std::uint32_t g = gen;
|
||||
this->in->async_read(round, inbox[round].data(), bytes,
|
||||
[self, g, round, bytes](const std::error_code & ec) {
|
||||
std::lock_guard<std::mutex> lock(self->mu);
|
||||
if (self->dead || g != self->gen)
|
||||
return;
|
||||
if (ec)
|
||||
{
|
||||
self->fail_locked(ec, "read round "
|
||||
+ std::to_string(round), false);
|
||||
return;
|
||||
}
|
||||
self->in_filled[round] = bytes;
|
||||
self->st.bytes_received += bytes;
|
||||
self->progress.fetch_add(bytes, std::memory_order_relaxed);
|
||||
});
|
||||
}
|
||||
|
||||
void read_lane_hdr_locked(std::size_t lane)
|
||||
{
|
||||
auto self = shared_from_this();
|
||||
const std::uint32_t g = gen;
|
||||
auto buf = hdr_bufs[lane];
|
||||
in->async_read(lane, buf->data(), buf->size(),
|
||||
[self, g, lane, buf](const std::error_code & ec) {
|
||||
self->on_lane_hdr(g, lane, buf, ec);
|
||||
});
|
||||
}
|
||||
|
||||
void on_lane_hdr(std::uint32_t g, std::size_t lane,
|
||||
const std::shared_ptr<std::array<std::uint8_t, round_lane_hdr::size>> & buf,
|
||||
const std::error_code & ec)
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(mu);
|
||||
if (dead || g != gen)
|
||||
return;
|
||||
if (ec)
|
||||
{
|
||||
// A peer that closes after its last frame finished cleanly.
|
||||
if (!complete_locked())
|
||||
fail_locked(ec, "read lane " + std::to_string(lane), false);
|
||||
return;
|
||||
}
|
||||
const auto hdr = round_lane_hdr::unpack(buf->data());
|
||||
if (hdr.round >= slots.size() || map.lane(hdr.round) != lane)
|
||||
{
|
||||
fail_locked(std::make_error_code(std::errc::protocol_error),
|
||||
"frame for round " + std::to_string(hdr.round) + " on lane "
|
||||
+ std::to_string(lane),
|
||||
true);
|
||||
return;
|
||||
}
|
||||
const std::size_t sb = slots[hdr.round];
|
||||
const std::size_t cap = count * sb;
|
||||
if (hdr.nbytes == 0)
|
||||
{
|
||||
if (sb == 0 && !zero_seen[hdr.round])
|
||||
{
|
||||
zero_seen[hdr.round] = true;
|
||||
progress.fetch_add(1, std::memory_order_relaxed);
|
||||
}
|
||||
read_lane_hdr_locked(lane);
|
||||
return;
|
||||
}
|
||||
if (sb == 0 || hdr.nbytes % sb != 0
|
||||
|| in_filled[hdr.round] + hdr.nbytes > cap)
|
||||
{
|
||||
fail_locked(std::make_error_code(std::errc::message_size),
|
||||
"round " + std::to_string(hdr.round) + " frame of "
|
||||
+ std::to_string(hdr.nbytes) + " bytes does not fit",
|
||||
true);
|
||||
return;
|
||||
}
|
||||
const std::uint16_t round = hdr.round;
|
||||
const std::size_t nbytes = hdr.nbytes;
|
||||
auto self = shared_from_this();
|
||||
in->async_read(lane, inbox[round].data() + in_filled[round], nbytes,
|
||||
[self, g, lane, round, nbytes](const std::error_code & xec) {
|
||||
std::lock_guard<std::mutex> lk(self->mu);
|
||||
if (self->dead || g != self->gen)
|
||||
return;
|
||||
if (xec)
|
||||
{
|
||||
self->fail_locked(xec, "read round "
|
||||
+ std::to_string(round), false);
|
||||
return;
|
||||
}
|
||||
self->in_filled[round] += nbytes;
|
||||
self->st.bytes_received += nbytes;
|
||||
self->progress.fetch_add(nbytes, std::memory_order_relaxed);
|
||||
self->read_lane_hdr_locked(lane);
|
||||
});
|
||||
}
|
||||
|
||||
// --- writes -----------------------------------------------------
|
||||
round_window & window(std::uint16_t round)
|
||||
{
|
||||
if (round >= win.size())
|
||||
throw std::out_of_range("async_round_sink round");
|
||||
return win[round];
|
||||
}
|
||||
|
||||
std::size_t flush_one(std::uint16_t round)
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(mu);
|
||||
auto & w = window(round);
|
||||
std::size_t begin = 0;
|
||||
std::size_t nslots = 0;
|
||||
const std::uint8_t * pend = w.pending_out(begin, nslots);
|
||||
const std::size_t sb = slots[round];
|
||||
const std::size_t nbytes = nslots * sb;
|
||||
const std::size_t lane = map.lane(round);
|
||||
auto self = shared_from_this();
|
||||
const std::uint32_t g = gen;
|
||||
auto done = [self, g](const std::error_code & ec) {
|
||||
self->on_write(g, ec);
|
||||
};
|
||||
if (map.framed)
|
||||
{
|
||||
if (nslots != 0 || (sb == 0 && !announced[round]))
|
||||
{
|
||||
auto frame = acquire_buffer(round_lane_hdr::size + nbytes);
|
||||
round_lane_hdr{round, static_cast<std::uint32_t>(nbytes)}.pack(
|
||||
frame->data());
|
||||
if (nbytes != 0)
|
||||
std::memcpy(frame->data() + round_lane_hdr::size, pend,
|
||||
nbytes);
|
||||
streams->async_write_owned(lane, std::move(frame), done);
|
||||
if (sb == 0)
|
||||
announced[round] = true;
|
||||
++st.flushes;
|
||||
st.bytes_sent += nbytes;
|
||||
}
|
||||
if (nslots != 0)
|
||||
w.mark_flushed(nslots);
|
||||
}
|
||||
else
|
||||
{
|
||||
if (nslots != 0 && sb != 0)
|
||||
{
|
||||
streams->async_write(lane, pend, nbytes, done);
|
||||
++st.flushes;
|
||||
st.bytes_sent += nbytes;
|
||||
}
|
||||
if (nslots != 0)
|
||||
w.mark_flushed(nslots);
|
||||
if (!legacy_started[round])
|
||||
{
|
||||
if (hello_ok)
|
||||
issue_legacy_locked(round);
|
||||
else
|
||||
legacy_wanted[round] = true;
|
||||
}
|
||||
}
|
||||
return lane;
|
||||
}
|
||||
|
||||
/// @brief Stall the producer while lane `lane`'s window is full.
|
||||
void await_window(std::size_t lane)
|
||||
{
|
||||
const std::size_t w = streams->lane_window_bytes(lane);
|
||||
if (w == 0 || streams->lane_buffered_bytes(lane) <= w)
|
||||
return;
|
||||
const auto t0 = std::chrono::steady_clock::now();
|
||||
auto since = t0;
|
||||
std::size_t last = streams->lane_buffered_bytes(lane);
|
||||
for (;;)
|
||||
{
|
||||
raise_if_failed();
|
||||
const std::size_t now_buf = streams->lane_buffered_bytes(lane);
|
||||
if (now_buf <= streams->lane_window_bytes(lane))
|
||||
break;
|
||||
const auto now = std::chrono::steady_clock::now();
|
||||
if (now_buf < last)
|
||||
{
|
||||
last = now_buf;
|
||||
since = now;
|
||||
}
|
||||
else if (now - since > opt.drain_timeout)
|
||||
{
|
||||
throw std::runtime_error("async_round_sink: lane "
|
||||
+ std::to_string(lane) + " did not drain for "
|
||||
+ std::to_string(opt.drain_timeout.count()) + " ms ("
|
||||
+ std::to_string(now_buf) + " bytes buffered, window "
|
||||
+ std::to_string(w) + ")");
|
||||
}
|
||||
restart_if_stopped();
|
||||
io->run_one_for(std::chrono::milliseconds(10));
|
||||
}
|
||||
const auto dt = std::chrono::steady_clock::now() - t0;
|
||||
std::lock_guard<std::mutex> lock(mu);
|
||||
++st.window_waits;
|
||||
st.window_wait_ns += static_cast<std::uint64_t>(
|
||||
std::chrono::duration_cast<std::chrono::nanoseconds>(dt).count());
|
||||
}
|
||||
|
||||
async_stream_array * streams = nullptr;
|
||||
async_stream_array * in = nullptr;
|
||||
asio::io_context * io = nullptr;
|
||||
std::size_t count = 1;
|
||||
std::vector<std::size_t> slots;
|
||||
round_lane_map map;
|
||||
sink_options opt;
|
||||
std::uint64_t fingerprint = 0;
|
||||
|
||||
std::mutex mu;
|
||||
std::vector<round_window> win;
|
||||
std::vector<std::vector<std::uint8_t>> inbox;
|
||||
std::vector<std::size_t> in_filled;
|
||||
std::vector<bool> zero_seen;
|
||||
std::vector<bool> announced;
|
||||
std::vector<bool> legacy_started;
|
||||
std::vector<bool> legacy_wanted;
|
||||
std::vector<std::shared_ptr<std::array<std::uint8_t, round_lane_hdr::size>>>
|
||||
hdr_bufs;
|
||||
std::shared_ptr<std::vector<std::uint8_t>> hello_in;
|
||||
std::uint32_t gen = 0;
|
||||
bool hello_ok = false;
|
||||
std::error_code fail;
|
||||
std::string fail_what;
|
||||
bool fatal = false;
|
||||
unsigned reconnects = 0;
|
||||
bool dead = false;
|
||||
std::atomic<std::uint64_t> progress{0};
|
||||
sink_stats st;
|
||||
};
|
||||
|
||||
std::shared_ptr<core> core_;
|
||||
};
|
||||
|
||||
} // namespace net
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_NET_ASYNC_ROUND_SINK_HPP__
|
||||
773
include/dpf/net/async_sctp_stream_array.hpp
Normal file
773
include/dpf/net/async_sctp_stream_array.hpp
Normal file
|
|
@ -0,0 +1,773 @@
|
|||
/// @file dpf/net/async_sctp_stream_array.hpp
|
||||
/// @brief Truly asynchronous SCTP-backed indexed byte streams (Linux/libsctp).
|
||||
/// @details One SCTP association carries every lane: index `i` maps to SCTP
|
||||
/// stream `i`. Receive is head-of-line free across streams. Writes are
|
||||
/// split into `wire_policy::chunk_bytes` messages and sent round-robin
|
||||
/// across streams, so a large write on one stream does not hold small
|
||||
/// writes on another behind it. The window covers the association.
|
||||
///
|
||||
/// Platform contract:
|
||||
/// * Linux with `<netinet/sctp.h>` (libsctp) → real backend, needs `-lsctp`.
|
||||
/// * Everything else → the class exists, every constructor throws.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_NET_ASYNC_SCTP_STREAM_ARRAY_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_NET_ASYNC_SCTP_STREAM_ARRAY_HPP__
|
||||
|
||||
#if !defined(DPF_HAS_LIBSCTP)
|
||||
# if defined(__linux__) && defined(__has_include)
|
||||
# if __has_include(<netinet/sctp.h>)
|
||||
# define DPF_HAS_LIBSCTP 1
|
||||
# else
|
||||
# define DPF_HAS_LIBSCTP 0
|
||||
# endif
|
||||
# else
|
||||
# define DPF_HAS_LIBSCTP 0
|
||||
# endif
|
||||
#endif
|
||||
|
||||
#include <atomic>
|
||||
#include <chrono>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <deque>
|
||||
#include <exception>
|
||||
#include <functional>
|
||||
#include <memory>
|
||||
#include <mutex>
|
||||
#include <stdexcept>
|
||||
#include <string>
|
||||
#include <system_error>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/net/asio_ns.hpp"
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
#include "dpf/net/async_stream_array.hpp"
|
||||
#include "dpf/net/connect.hpp"
|
||||
#include "dpf/net/policy.hpp"
|
||||
|
||||
#if DPF_HAS_LIBSCTP
|
||||
# include <arpa/inet.h>
|
||||
# include <fcntl.h>
|
||||
# include <netinet/in.h>
|
||||
# include <netinet/sctp.h>
|
||||
# include <sys/socket.h>
|
||||
# include <unistd.h>
|
||||
# include <cerrno>
|
||||
#endif
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace net
|
||||
{
|
||||
|
||||
/// @brief True when this build has a real SCTP backend.
|
||||
inline constexpr bool sctp_available() noexcept
|
||||
{
|
||||
return DPF_HAS_LIBSCTP != 0;
|
||||
}
|
||||
|
||||
#if DPF_HAS_LIBSCTP
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
/// @brief Stream counts, per-message `sinfo`, and socket options.
|
||||
inline void sctp_configure(int fd, std::size_t nstreams,
|
||||
const socket_options & o = {})
|
||||
{
|
||||
struct sctp_initmsg im;
|
||||
std::memset(&im, 0, sizeof(im));
|
||||
im.sinit_num_ostreams = static_cast<std::uint16_t>(nstreams);
|
||||
im.sinit_max_instreams = static_cast<std::uint16_t>(nstreams);
|
||||
(void)::setsockopt(fd, IPPROTO_SCTP, SCTP_INITMSG, &im, sizeof(im));
|
||||
|
||||
struct sctp_event_subscribe ev;
|
||||
std::memset(&ev, 0, sizeof(ev));
|
||||
ev.sctp_data_io_event = 1;
|
||||
(void)::setsockopt(fd, IPPROTO_SCTP, SCTP_EVENTS, &ev, sizeof(ev));
|
||||
|
||||
const int nodelay = o.no_delay ? 1 : 0;
|
||||
(void)::setsockopt(fd, IPPROTO_SCTP, SCTP_NODELAY, &nodelay, sizeof(nodelay));
|
||||
if (o.send_buffer > 0)
|
||||
(void)::setsockopt(fd, SOL_SOCKET, SO_SNDBUF, &o.send_buffer,
|
||||
sizeof(o.send_buffer));
|
||||
if (o.recv_buffer > 0)
|
||||
(void)::setsockopt(fd, SOL_SOCKET, SO_RCVBUF, &o.recv_buffer,
|
||||
sizeof(o.recv_buffer));
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/// @brief N lanes over one SCTP association, fully overlapped.
|
||||
class async_sctp_stream_array final : public async_stream_array
|
||||
{
|
||||
public:
|
||||
/// @brief Adopt a connected one-to-one SCTP fd.
|
||||
async_sctp_stream_array(asio::io_context & io, int fd, std::size_t nstreams,
|
||||
const wire_policy & pol = {})
|
||||
{
|
||||
if (nstreams == 0)
|
||||
throw std::invalid_argument("async_sctp_stream_array empty");
|
||||
if (nstreams > 0xffffu)
|
||||
throw std::invalid_argument("async_sctp_stream_array: stream ids are u16");
|
||||
pol.validate();
|
||||
detail::sctp_configure(fd, nstreams, pol.socket);
|
||||
impl_ = std::make_shared<impl>(io, fd, nstreams, pol);
|
||||
impl_->start();
|
||||
}
|
||||
|
||||
async_sctp_stream_array(const async_sctp_stream_array &) = delete;
|
||||
async_sctp_stream_array & operator=(const async_sctp_stream_array &) = delete;
|
||||
|
||||
~async_sctp_stream_array() override
|
||||
{
|
||||
if (impl_)
|
||||
impl_->close_graceful();
|
||||
}
|
||||
|
||||
std::size_t size() const noexcept override { return impl_->n; }
|
||||
asio::io_context & context() noexcept override { return impl_->io; }
|
||||
|
||||
void async_write(std::size_t i, const void * src, std::size_t n,
|
||||
async_handler h) override
|
||||
{
|
||||
impl_->write(i, n == 0 ? nullptr : copy_buffer(src, n), std::move(h));
|
||||
}
|
||||
|
||||
void async_write_owned(std::size_t i,
|
||||
std::shared_ptr<std::vector<std::uint8_t>> buf, async_handler h) override
|
||||
{
|
||||
impl_->write(i, std::move(buf), std::move(h));
|
||||
}
|
||||
|
||||
void async_read(std::size_t i, void * dst, std::size_t n,
|
||||
async_handler h) override
|
||||
{
|
||||
impl_->read(i, dst, n, std::move(h));
|
||||
}
|
||||
|
||||
std::size_t buffered_bytes() const noexcept override
|
||||
{
|
||||
return impl_->buffered.load(std::memory_order_relaxed);
|
||||
}
|
||||
|
||||
std::size_t window_bytes() const noexcept override
|
||||
{
|
||||
return impl_->window.load(std::memory_order_relaxed);
|
||||
}
|
||||
|
||||
void set_window_bytes(std::size_t bytes) override
|
||||
{
|
||||
impl_->window.store(bytes, std::memory_order_relaxed);
|
||||
}
|
||||
|
||||
stream_stats stats() const override { return impl_->stats(); }
|
||||
|
||||
void close() noexcept override
|
||||
{
|
||||
impl_->close(asio::error::operation_aborted);
|
||||
}
|
||||
|
||||
private:
|
||||
struct chunk
|
||||
{
|
||||
std::shared_ptr<std::vector<std::uint8_t>> payload;
|
||||
std::size_t off = 0;
|
||||
std::size_t len = 0;
|
||||
std::shared_ptr<detail::write_op> op;
|
||||
};
|
||||
|
||||
struct impl : std::enable_shared_from_this<impl>
|
||||
{
|
||||
impl(asio::io_context & io_, int fd, std::size_t nstreams,
|
||||
const wire_policy & pol_)
|
||||
: io(io_),
|
||||
sd(io_),
|
||||
strand(asio::make_strand(io_)),
|
||||
n(nstreams),
|
||||
pol(pol_),
|
||||
inbox(nstreams),
|
||||
partial(nstreams),
|
||||
wait(nstreams),
|
||||
outq(nstreams),
|
||||
rbuf(std::max<std::size_t>(pol_.chunk_bytes, std::size_t{1} << 16)),
|
||||
window(pol_.window_bytes)
|
||||
{
|
||||
const int fl = ::fcntl(fd, F_GETFL, 0);
|
||||
if (fl >= 0)
|
||||
(void)::fcntl(fd, F_SETFL, fl | O_NONBLOCK);
|
||||
sd.assign(fd);
|
||||
}
|
||||
|
||||
void check(std::size_t i) const
|
||||
{
|
||||
if (i >= n)
|
||||
throw std::out_of_range("async_sctp_stream_array index "
|
||||
+ std::to_string(i));
|
||||
}
|
||||
|
||||
void start()
|
||||
{
|
||||
auto self = shared_from_this();
|
||||
asio::post(strand, [self] { self->arm_read(); });
|
||||
}
|
||||
|
||||
// --- reads ---------------------------------------------------------
|
||||
void read(std::size_t i, void * dst, std::size_t nbytes, async_handler h)
|
||||
{
|
||||
check(i);
|
||||
std::lock_guard<std::mutex> lock(mu);
|
||||
if (fail)
|
||||
{
|
||||
detail::post_handler(io, std::move(h), fail);
|
||||
return;
|
||||
}
|
||||
auto & w = wait[i];
|
||||
if (w.active)
|
||||
throw std::logic_error(
|
||||
"async_sctp_stream_array: overlapping read on stream "
|
||||
+ std::to_string(i));
|
||||
const std::size_t k = inbox[i].take(dst, nbytes, pol.compact_bytes);
|
||||
if (k == nbytes)
|
||||
{
|
||||
detail::post_handler(io, std::move(h), {});
|
||||
return;
|
||||
}
|
||||
w.active = true;
|
||||
w.dst = dst;
|
||||
w.n = nbytes;
|
||||
w.filled = k;
|
||||
w.h = std::move(h);
|
||||
}
|
||||
|
||||
void satisfy_locked(std::size_t i)
|
||||
{
|
||||
auto & w = wait[i];
|
||||
if (!w.active)
|
||||
return;
|
||||
w.filled += inbox[i].take(static_cast<std::uint8_t *>(w.dst) + w.filled,
|
||||
w.n - w.filled, pol.compact_bytes);
|
||||
if (w.filled == w.n)
|
||||
{
|
||||
w.active = false;
|
||||
detail::post_handler(io, std::move(w.h), {});
|
||||
}
|
||||
}
|
||||
|
||||
void arm_read()
|
||||
{
|
||||
auto self = shared_from_this();
|
||||
sd.async_wait(asio::posix::stream_descriptor::wait_read,
|
||||
asio::bind_executor(strand, [self](const std::error_code & ec) {
|
||||
self->on_readable(ec);
|
||||
}));
|
||||
}
|
||||
|
||||
void on_readable(const std::error_code & ec)
|
||||
{
|
||||
if (ec)
|
||||
{
|
||||
deliver_error(ec);
|
||||
return;
|
||||
}
|
||||
for (;;)
|
||||
{
|
||||
struct sctp_sndrcvinfo sinfo;
|
||||
std::memset(&sinfo, 0, sizeof(sinfo));
|
||||
int flags = 0;
|
||||
const ssize_t r = ::sctp_recvmsg(sd.native_handle(), rbuf.data(),
|
||||
rbuf.size(), nullptr, nullptr, &sinfo, &flags);
|
||||
if (r < 0)
|
||||
{
|
||||
if (errno == EAGAIN || errno == EWOULDBLOCK)
|
||||
break;
|
||||
deliver_error(std::error_code(errno, std::generic_category()));
|
||||
return;
|
||||
}
|
||||
if (r == 0)
|
||||
{
|
||||
deliver_error(asio::error::eof);
|
||||
return;
|
||||
}
|
||||
if (flags & MSG_NOTIFICATION)
|
||||
continue;
|
||||
const std::size_t stream = sinfo.sinfo_stream;
|
||||
if (stream >= n)
|
||||
{
|
||||
deliver_error(std::make_error_code(std::errc::protocol_error));
|
||||
return;
|
||||
}
|
||||
bool too_big = false;
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(mu);
|
||||
if (closed)
|
||||
return;
|
||||
auto & part = partial[stream];
|
||||
part.insert(part.end(), rbuf.begin(), rbuf.begin() + r);
|
||||
if (part.size() > pol.max_frame)
|
||||
too_big = true;
|
||||
else if ((flags & MSG_EOR) != 0)
|
||||
{
|
||||
counters.read(part.size(), part.size(), 1);
|
||||
auto & box = inbox[stream].bytes;
|
||||
box.insert(box.end(), part.begin(), part.end());
|
||||
part.clear();
|
||||
satisfy_locked(stream);
|
||||
}
|
||||
}
|
||||
if (too_big)
|
||||
{
|
||||
deliver_error(std::make_error_code(std::errc::message_size));
|
||||
return;
|
||||
}
|
||||
}
|
||||
arm_read();
|
||||
}
|
||||
|
||||
void fail_waiters_locked(const std::error_code & ec)
|
||||
{
|
||||
for (auto & w : wait)
|
||||
{
|
||||
if (!w.active)
|
||||
continue;
|
||||
w.active = false;
|
||||
detail::post_handler(io, std::move(w.h), ec);
|
||||
}
|
||||
}
|
||||
|
||||
void deliver_error(const std::error_code & ec)
|
||||
{
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(mu);
|
||||
if (!fail)
|
||||
fail = ec;
|
||||
fail_waiters_locked(fail);
|
||||
}
|
||||
auto self = shared_from_this();
|
||||
asio::post(strand, [self, ec] { self->fail_queued(ec); });
|
||||
}
|
||||
|
||||
void close(const std::error_code & ec)
|
||||
{
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(mu);
|
||||
if (!fail)
|
||||
fail = ec;
|
||||
if (closed && aborted)
|
||||
return;
|
||||
closed = true;
|
||||
aborted = true;
|
||||
fail_waiters_locked(ec);
|
||||
}
|
||||
auto self = shared_from_this();
|
||||
asio::post(strand, [self, ec] {
|
||||
std::error_code e;
|
||||
self->sd.close(e);
|
||||
self->fail_queued(ec);
|
||||
});
|
||||
}
|
||||
|
||||
void close_graceful()
|
||||
{
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(mu);
|
||||
if (closed)
|
||||
return;
|
||||
closed = true;
|
||||
fail_waiters_locked(asio::error::operation_aborted);
|
||||
}
|
||||
auto self = shared_from_this();
|
||||
asio::post(strand, [self] {
|
||||
self->draining = true;
|
||||
if (!self->writing)
|
||||
{
|
||||
std::error_code e;
|
||||
self->sd.close(e);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
stream_stats stats() const
|
||||
{
|
||||
stream_stats s;
|
||||
counters.fill(s);
|
||||
s.buffered = buffered.load(std::memory_order_relaxed);
|
||||
std::lock_guard<std::mutex> lock(mu);
|
||||
for (const auto & b : inbox)
|
||||
s.unread += b.avail();
|
||||
s.error = fail;
|
||||
s.closed = closed;
|
||||
return s;
|
||||
}
|
||||
|
||||
// --- writes (strand) ----------------------------------------------
|
||||
void write(std::size_t i, std::shared_ptr<std::vector<std::uint8_t>> payload,
|
||||
async_handler h)
|
||||
{
|
||||
check(i);
|
||||
const std::size_t len = payload ? payload->size() : 0;
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(mu);
|
||||
if (fail || closed)
|
||||
{
|
||||
detail::post_handler(io, std::move(h),
|
||||
fail ? fail : asio::error::operation_aborted);
|
||||
return;
|
||||
}
|
||||
}
|
||||
if (len == 0)
|
||||
{
|
||||
detail::post_handler(io, std::move(h), {});
|
||||
return;
|
||||
}
|
||||
const auto pieces = detail::split_chunks(len, pol.chunk_bytes);
|
||||
auto op = std::make_shared<detail::write_op>();
|
||||
op->h = std::move(h);
|
||||
op->left = pieces.size();
|
||||
std::vector<chunk> cs;
|
||||
cs.reserve(pieces.size());
|
||||
for (const auto & pc : pieces)
|
||||
cs.push_back(chunk{payload, pc.first, pc.second, op});
|
||||
buffered.fetch_add(len, std::memory_order_relaxed);
|
||||
auto self = shared_from_this();
|
||||
asio::post(strand, [self, i, cs = std::move(cs)]() mutable {
|
||||
std::error_code ec;
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(self->mu);
|
||||
ec = self->fail;
|
||||
}
|
||||
if (ec)
|
||||
{
|
||||
for (auto & c : cs)
|
||||
self->finish_chunk(c, ec);
|
||||
return;
|
||||
}
|
||||
for (auto & c : cs)
|
||||
self->outq[i].push_back(std::move(c));
|
||||
if (!self->writing)
|
||||
self->write_next();
|
||||
});
|
||||
}
|
||||
|
||||
void finish_chunk(chunk & c, const std::error_code & ec)
|
||||
{
|
||||
buffered.fetch_sub(std::min(buffered.load(std::memory_order_relaxed),
|
||||
c.len),
|
||||
std::memory_order_relaxed);
|
||||
auto & op = *c.op;
|
||||
if (ec && !op.ec)
|
||||
op.ec = ec;
|
||||
if (--op.left == 0)
|
||||
detail::post_handler(io, std::move(op.h), op.ec);
|
||||
}
|
||||
|
||||
bool pick(std::size_t & lane)
|
||||
{
|
||||
for (std::size_t k = 0; k < n; ++k)
|
||||
{
|
||||
const std::size_t l = (rr + k) % n;
|
||||
if (!outq[l].empty())
|
||||
{
|
||||
lane = l;
|
||||
return true;
|
||||
}
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
void write_next()
|
||||
{
|
||||
std::size_t lane = 0;
|
||||
if (!pick(lane))
|
||||
{
|
||||
writing = false;
|
||||
if (draining)
|
||||
{
|
||||
std::error_code e;
|
||||
sd.close(e);
|
||||
}
|
||||
return;
|
||||
}
|
||||
writing = true;
|
||||
auto & front = outq[lane].front();
|
||||
const ssize_t r = ::sctp_sendmsg(sd.native_handle(),
|
||||
front.payload->data() + front.off, front.len, nullptr, 0, 0, 0,
|
||||
static_cast<std::uint16_t>(lane), 0, 0);
|
||||
if (r < 0)
|
||||
{
|
||||
if (errno == EAGAIN || errno == EWOULDBLOCK)
|
||||
{
|
||||
auto self = shared_from_this();
|
||||
sd.async_wait(asio::posix::stream_descriptor::wait_write,
|
||||
asio::bind_executor(strand,
|
||||
[self](const std::error_code & ec) {
|
||||
if (ec)
|
||||
{
|
||||
self->writing = false;
|
||||
self->deliver_error(ec);
|
||||
return;
|
||||
}
|
||||
self->write_next();
|
||||
}));
|
||||
return;
|
||||
}
|
||||
writing = false;
|
||||
deliver_error(std::error_code(errno, std::generic_category()));
|
||||
return;
|
||||
}
|
||||
const std::size_t sent = static_cast<std::size_t>(r);
|
||||
counters.wrote(sent, sent, 1);
|
||||
if (sent < front.len)
|
||||
{
|
||||
buffered.fetch_sub(std::min(buffered.load(), sent),
|
||||
std::memory_order_relaxed);
|
||||
front.off += sent;
|
||||
front.len -= sent;
|
||||
write_next();
|
||||
return;
|
||||
}
|
||||
chunk done = std::move(front);
|
||||
outq[lane].pop_front();
|
||||
rr = (lane + 1) % n;
|
||||
finish_chunk(done, {});
|
||||
auto self = shared_from_this();
|
||||
asio::post(strand, [self] { self->write_next(); });
|
||||
}
|
||||
|
||||
void fail_queued(const std::error_code & ec)
|
||||
{
|
||||
writing = false;
|
||||
for (auto & q : outq)
|
||||
{
|
||||
for (auto & c : q)
|
||||
finish_chunk(c, ec);
|
||||
q.clear();
|
||||
}
|
||||
}
|
||||
|
||||
asio::io_context & io;
|
||||
asio::posix::stream_descriptor sd;
|
||||
asio::strand<asio::io_context::executor_type> strand;
|
||||
std::size_t n = 0;
|
||||
wire_policy pol;
|
||||
mutable std::mutex mu;
|
||||
std::vector<detail::stream_inbox> inbox;
|
||||
std::vector<std::vector<std::uint8_t>> partial;
|
||||
std::vector<detail::stream_waiter> wait;
|
||||
std::error_code fail;
|
||||
bool closed = false;
|
||||
bool aborted = false;
|
||||
std::vector<std::deque<chunk>> outq;
|
||||
std::size_t rr = 0;
|
||||
bool writing = false;
|
||||
bool draining = false;
|
||||
std::vector<std::uint8_t> rbuf;
|
||||
std::atomic<std::size_t> buffered{0};
|
||||
std::atomic<std::size_t> window;
|
||||
detail::io_counters counters;
|
||||
};
|
||||
|
||||
std::shared_ptr<impl> impl_;
|
||||
};
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Association setup (deadline-bounded)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// @brief Listening one-to-one SCTP socket.
|
||||
class sctp_listener
|
||||
{
|
||||
public:
|
||||
sctp_listener(unsigned short port, std::size_t nstreams,
|
||||
const socket_options & o = {})
|
||||
: nstreams_(nstreams), opts_(o)
|
||||
{
|
||||
if (nstreams == 0)
|
||||
throw std::invalid_argument("sctp_listener: no streams");
|
||||
fd_ = ::socket(AF_INET, SOCK_STREAM, IPPROTO_SCTP);
|
||||
if (fd_ < 0)
|
||||
throw std::system_error(errno, std::generic_category(),
|
||||
"sctp_listener: socket");
|
||||
int one = 1;
|
||||
(void)::setsockopt(fd_, SOL_SOCKET, SO_REUSEADDR, &one, sizeof(one));
|
||||
detail::sctp_configure(fd_, nstreams, o);
|
||||
struct sockaddr_in addr;
|
||||
std::memset(&addr, 0, sizeof(addr));
|
||||
addr.sin_family = AF_INET;
|
||||
addr.sin_addr.s_addr = htonl(INADDR_ANY);
|
||||
addr.sin_port = htons(port);
|
||||
if (::bind(fd_, reinterpret_cast<struct sockaddr *>(&addr), sizeof(addr)) < 0
|
||||
|| ::listen(fd_, 16) < 0)
|
||||
{
|
||||
const int e = errno;
|
||||
::close(fd_);
|
||||
fd_ = -1;
|
||||
throw std::system_error(e, std::generic_category(),
|
||||
"sctp_listener: bind/listen on port " + std::to_string(port));
|
||||
}
|
||||
socklen_t alen = sizeof(addr);
|
||||
if (::getsockname(fd_, reinterpret_cast<struct sockaddr *>(&addr), &alen) == 0)
|
||||
port_ = ntohs(addr.sin_port);
|
||||
}
|
||||
|
||||
sctp_listener(const sctp_listener &) = delete;
|
||||
sctp_listener & operator=(const sctp_listener &) = delete;
|
||||
|
||||
~sctp_listener()
|
||||
{
|
||||
if (fd_ >= 0)
|
||||
::close(fd_);
|
||||
}
|
||||
|
||||
unsigned short port() const noexcept { return port_; }
|
||||
int native_handle() const noexcept { return fd_; }
|
||||
|
||||
/// @brief Accept one association within `budget`; returns the fd.
|
||||
int accept(std::chrono::milliseconds budget)
|
||||
{
|
||||
const auto deadline = setup_clock::now() + budget;
|
||||
detail::wait_fd(fd_, POLLIN, deadline,
|
||||
"sctp accept on port " + std::to_string(port_));
|
||||
const int cfd = ::accept(fd_, nullptr, nullptr);
|
||||
if (cfd < 0)
|
||||
throw std::system_error(errno, std::generic_category(), "sctp accept");
|
||||
detail::sctp_configure(cfd, nstreams_, opts_);
|
||||
return cfd;
|
||||
}
|
||||
|
||||
private:
|
||||
int fd_ = -1;
|
||||
unsigned short port_ = 0;
|
||||
std::size_t nstreams_ = 0;
|
||||
socket_options opts_{};
|
||||
};
|
||||
|
||||
/// @brief Connect one SCTP association, retrying refusals until `budget`.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline int connect_sctp(const std::string & host, unsigned short port,
|
||||
std::size_t nstreams, std::chrono::milliseconds budget,
|
||||
const socket_options & o = {})
|
||||
{
|
||||
const auto deadline = setup_clock::now() + budget;
|
||||
struct sockaddr_in addr;
|
||||
std::memset(&addr, 0, sizeof(addr));
|
||||
addr.sin_family = AF_INET;
|
||||
addr.sin_port = htons(port);
|
||||
const std::string h = host == "localhost" ? "127.0.0.1" : host;
|
||||
if (::inet_pton(AF_INET, h.c_str(), &addr.sin_addr) != 1)
|
||||
throw std::invalid_argument("connect_sctp: bad host " + host);
|
||||
int last = ETIMEDOUT;
|
||||
for (;;)
|
||||
{
|
||||
const int fd = ::socket(AF_INET, SOCK_STREAM, IPPROTO_SCTP);
|
||||
if (fd < 0)
|
||||
throw std::system_error(errno, std::generic_category(),
|
||||
"connect_sctp: socket");
|
||||
detail::sctp_configure(fd, nstreams, o);
|
||||
if (::connect(fd, reinterpret_cast<struct sockaddr *>(&addr),
|
||||
sizeof(addr)) == 0)
|
||||
return fd;
|
||||
last = errno;
|
||||
::close(fd);
|
||||
if (setup_clock::now() >= deadline)
|
||||
throw std::system_error(last, std::generic_category(),
|
||||
"connect_sctp " + host + ":" + std::to_string(port));
|
||||
std::this_thread::sleep_for(std::chrono::milliseconds(20));
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Accept one association on an ephemeral (or given) port.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline int accept_sctp_association(asio::io_context &,
|
||||
std::atomic<unsigned short> & port, std::size_t nstreams,
|
||||
std::chrono::milliseconds budget = deadlines{}.accept)
|
||||
{
|
||||
sctp_listener lst(port.load(), nstreams);
|
||||
port.store(lst.port());
|
||||
return lst.accept(budget);
|
||||
}
|
||||
|
||||
/// @brief Connect one association to `host:port`.
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline int connect_sctp_association(asio::io_context &, const std::string & host,
|
||||
std::atomic<unsigned short> & port, std::size_t nstreams,
|
||||
std::chrono::milliseconds budget = deadlines{}.connect)
|
||||
{
|
||||
return connect_sctp(host, port.load(), nstreams, budget);
|
||||
}
|
||||
|
||||
#else // !DPF_HAS_LIBSCTP
|
||||
|
||||
class async_sctp_stream_array final : public async_stream_array
|
||||
{
|
||||
public:
|
||||
async_sctp_stream_array(asio::io_context &, int, std::size_t,
|
||||
const wire_policy & = {})
|
||||
{
|
||||
throw std::logic_error(
|
||||
"async_sctp_stream_array: real SCTP requires Linux + libsctp "
|
||||
"(<netinet/sctp.h>, link -lsctp)");
|
||||
}
|
||||
|
||||
std::size_t size() const noexcept override { return 0; }
|
||||
asio::io_context & context() noexcept override
|
||||
{
|
||||
std::terminate();
|
||||
}
|
||||
void async_write(std::size_t, const void *, std::size_t,
|
||||
async_handler) override
|
||||
{
|
||||
throw std::logic_error("async_sctp_stream_array: not available");
|
||||
}
|
||||
void async_read(std::size_t, void *, std::size_t, async_handler) override
|
||||
{
|
||||
throw std::logic_error("async_sctp_stream_array: not available");
|
||||
}
|
||||
};
|
||||
|
||||
class sctp_listener
|
||||
{
|
||||
public:
|
||||
sctp_listener(unsigned short, std::size_t, const socket_options & = {})
|
||||
{
|
||||
throw std::logic_error("sctp_listener: SCTP requires Linux + libsctp");
|
||||
}
|
||||
unsigned short port() const noexcept { return 0; }
|
||||
int native_handle() const noexcept { return -1; }
|
||||
int accept(std::chrono::milliseconds) { return -1; }
|
||||
};
|
||||
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline int connect_sctp(const std::string &, unsigned short, std::size_t,
|
||||
std::chrono::milliseconds, const socket_options & = {})
|
||||
{
|
||||
throw std::logic_error("connect_sctp: SCTP requires Linux + libsctp");
|
||||
}
|
||||
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline int accept_sctp_association(asio::io_context &,
|
||||
std::atomic<unsigned short> &, std::size_t,
|
||||
std::chrono::milliseconds = deadlines{}.accept)
|
||||
{
|
||||
throw std::logic_error(
|
||||
"accept_sctp_association: SCTP requires Linux + libsctp");
|
||||
}
|
||||
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
inline int connect_sctp_association(asio::io_context &, const std::string &,
|
||||
std::atomic<unsigned short> &, std::size_t,
|
||||
std::chrono::milliseconds = deadlines{}.connect)
|
||||
{
|
||||
throw std::logic_error(
|
||||
"connect_sctp_association: SCTP requires Linux + libsctp");
|
||||
}
|
||||
|
||||
#endif // DPF_HAS_LIBSCTP
|
||||
|
||||
} // namespace net
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_NET_ASYNC_SCTP_STREAM_ARRAY_HPP__
|
||||
2069
include/dpf/net/async_stream_array.hpp
Normal file
2069
include/dpf/net/async_stream_array.hpp
Normal file
File diff suppressed because it is too large
Load diff
98
include/dpf/net/buffer_pool.hpp
Normal file
98
include/dpf/net/buffer_pool.hpp
Normal file
|
|
@ -0,0 +1,98 @@
|
|||
/// @file dpf/net/buffer_pool.hpp
|
||||
/// @brief Recycled byte buffers for framed writes.
|
||||
/// @details `acquire_buffer` hands out a `shared_ptr` whose deleter returns the
|
||||
/// vector to a process-wide free list (capped). Callers that already
|
||||
/// own a buffer pass it to `async_write_owned` so the socket path can
|
||||
/// scatter-gather without a second payload copy.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_NET_BUFFER_POOL_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_NET_BUFFER_POOL_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <memory>
|
||||
#include <mutex>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/net/policy.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace net
|
||||
{
|
||||
|
||||
/// @brief Default outstanding-byte window (see `wire_policy::window_bytes`).
|
||||
inline std::size_t default_wire_window() noexcept
|
||||
{
|
||||
return wire_policy{}.window_bytes;
|
||||
}
|
||||
|
||||
/// @brief Default receive cap (see `wire_policy::max_frame`).
|
||||
inline constexpr std::size_t k_max_frame = std::size_t{16} << 20;
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
struct pool_state
|
||||
{
|
||||
std::mutex mu;
|
||||
std::vector<std::unique_ptr<std::vector<std::uint8_t>>> free;
|
||||
};
|
||||
|
||||
inline pool_state & buffers()
|
||||
{
|
||||
static pool_state s;
|
||||
return s;
|
||||
}
|
||||
|
||||
inline void recycle(std::vector<std::uint8_t> * p)
|
||||
{
|
||||
std::unique_ptr<std::vector<std::uint8_t>> owned(p);
|
||||
owned->clear();
|
||||
if (owned->capacity() > (std::size_t{1} << 20))
|
||||
return;
|
||||
auto & st = buffers();
|
||||
std::lock_guard<std::mutex> lock(st.mu);
|
||||
if (st.free.size() < 128)
|
||||
st.free.push_back(std::move(owned));
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/// @brief Buffer of at least `n` bytes, returned to the pool on last release.
|
||||
inline std::shared_ptr<std::vector<std::uint8_t>> acquire_buffer(std::size_t n)
|
||||
{
|
||||
std::unique_ptr<std::vector<std::uint8_t>> raw;
|
||||
{
|
||||
auto & st = detail::buffers();
|
||||
std::lock_guard<std::mutex> lock(st.mu);
|
||||
if (!st.free.empty())
|
||||
{
|
||||
raw = std::move(st.free.back());
|
||||
st.free.pop_back();
|
||||
}
|
||||
}
|
||||
if (!raw)
|
||||
raw.reset(new std::vector<std::uint8_t>());
|
||||
if (raw->capacity() < n)
|
||||
raw->reserve(n);
|
||||
raw->resize(n);
|
||||
std::vector<std::uint8_t> * p = raw.release();
|
||||
return std::shared_ptr<std::vector<std::uint8_t>>(p, &detail::recycle);
|
||||
}
|
||||
|
||||
/// @brief Copy `n` bytes into a pooled buffer.
|
||||
inline std::shared_ptr<std::vector<std::uint8_t>> copy_buffer(const void * src,
|
||||
std::size_t n)
|
||||
{
|
||||
auto buf = acquire_buffer(n);
|
||||
if (n != 0 && src != nullptr)
|
||||
std::memcpy(buf->data(), src, n);
|
||||
return buf;
|
||||
}
|
||||
|
||||
} // namespace net
|
||||
} // namespace dpf
|
||||
|
||||
#endif
|
||||
484
include/dpf/net/channel.hpp
Normal file
484
include/dpf/net/channel.hpp
Normal file
|
|
@ -0,0 +1,484 @@
|
|||
/// @file dpf/net/channel.hpp
|
||||
/// @brief Framed duplex stream for party message exchange.
|
||||
/// @details Every message is `u32` little-endian length, `u16` type tag, then
|
||||
/// payload. `exchange` orders send/recv by role so a single stream
|
||||
/// cannot deadlock. Call sites name `send` / `recv` / `exchange`;
|
||||
/// they do not touch ASIO buffers.
|
||||
/// @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_NET_CHANNEL_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_NET_CHANNEL_HPP__
|
||||
|
||||
#include <array>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <memory>
|
||||
#include <stdexcept>
|
||||
#include <string>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/net/asio_ns.hpp"
|
||||
#include "dpf/net/tls.hpp"
|
||||
|
||||
#include "hedley/hedley.h"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace net
|
||||
{
|
||||
|
||||
/// @brief Bytes and frames observed on one channel since the last reset.
|
||||
/// @details Counts include the 6-byte frame header. A failed read or write
|
||||
/// does not add to the tally. `payload_*` is the body alone;
|
||||
/// `exchanges` counts `exchange()` calls (one logical round each).
|
||||
struct io_tally
|
||||
{
|
||||
std::uint64_t bytes_sent = 0;
|
||||
std::uint64_t bytes_recv = 0;
|
||||
std::uint64_t frames_sent = 0;
|
||||
std::uint64_t frames_recv = 0;
|
||||
std::uint64_t payload_sent = 0;
|
||||
std::uint64_t payload_recv = 0;
|
||||
std::uint64_t exchanges = 0;
|
||||
|
||||
io_tally & operator+=(const io_tally & other) noexcept
|
||||
{
|
||||
bytes_sent += other.bytes_sent;
|
||||
bytes_recv += other.bytes_recv;
|
||||
frames_sent += other.frames_sent;
|
||||
frames_recv += other.frames_recv;
|
||||
payload_sent += other.payload_sent;
|
||||
payload_recv += other.payload_recv;
|
||||
exchanges += other.exchanges;
|
||||
return *this;
|
||||
}
|
||||
};
|
||||
|
||||
enum class msg : std::uint16_t
|
||||
{
|
||||
hangup = 0,
|
||||
beaver_tape = 1,
|
||||
ring_vector = 2,
|
||||
delta = 3,
|
||||
dpf_key = 4,
|
||||
proof_token = 5,
|
||||
case_ok = 6,
|
||||
case_fail = 7,
|
||||
bytes = 8,
|
||||
mac_key = 9,
|
||||
mac_share = 10,
|
||||
sketch_share = 11,
|
||||
round_batch = 12,
|
||||
};
|
||||
|
||||
inline constexpr std::uint16_t to_u16(msg t) noexcept
|
||||
{
|
||||
return static_cast<std::uint16_t>(t);
|
||||
}
|
||||
|
||||
/// @brief One framed duplex byte stream.
|
||||
class channel
|
||||
{
|
||||
public:
|
||||
using tcp_socket = asio::ip::tcp::socket;
|
||||
using local_socket = asio::local::stream_protocol::socket;
|
||||
|
||||
channel() = default;
|
||||
|
||||
explicit channel(tcp_socket sock)
|
||||
: tcp_(std::make_unique<tcp_socket>(std::move(sock)))
|
||||
{ }
|
||||
|
||||
explicit channel(local_socket sock)
|
||||
: local_(std::make_unique<local_socket>(std::move(sock)))
|
||||
{ }
|
||||
|
||||
#if DPF_HAS_OPENSSL
|
||||
/// @brief Framed channel over an established TLS 1.3 TCP stream.
|
||||
static channel from_tls(asio::io_context &, tls_stream s,
|
||||
std::shared_ptr<tls_context> ctx)
|
||||
{
|
||||
channel c;
|
||||
c.tls_ctx_ = std::move(ctx);
|
||||
c.tls_ = std::make_unique<tls_stream>(std::move(s));
|
||||
return c;
|
||||
}
|
||||
|
||||
/// @brief Framed channel over an established TLS 1.3 unix-domain stream.
|
||||
static channel from_tls_local(asio::io_context &, tls_local_stream s,
|
||||
std::shared_ptr<tls_context> ctx)
|
||||
{
|
||||
channel c;
|
||||
c.tls_ctx_ = std::move(ctx);
|
||||
c.tls_local_ = std::make_unique<tls_local_stream>(std::move(s));
|
||||
return c;
|
||||
}
|
||||
#endif
|
||||
|
||||
channel(channel &&) noexcept = default;
|
||||
channel & operator=(channel &&) noexcept = default;
|
||||
|
||||
channel(const channel &) = delete;
|
||||
channel & operator=(const channel &) = delete;
|
||||
|
||||
/// @brief Replay `frames` as a recv-only dealer tape. `send` throws.
|
||||
static channel from_inbox(std::vector<std::uint8_t> frames)
|
||||
{
|
||||
channel c;
|
||||
c.inbox_ = std::make_unique<inbox_buf>();
|
||||
c.inbox_->bytes = std::move(frames);
|
||||
return c;
|
||||
}
|
||||
|
||||
/// @brief True when an inbox tape has been read through its last byte.
|
||||
HEDLEY_NO_THROW
|
||||
bool inbox_done() const noexcept
|
||||
{
|
||||
return inbox_ != nullptr && inbox_->pos == inbox_->bytes.size();
|
||||
}
|
||||
|
||||
HEDLEY_NO_THROW
|
||||
bool open() const noexcept
|
||||
{
|
||||
return tcp_ != nullptr || local_ != nullptr || inbox_ != nullptr
|
||||
#if DPF_HAS_OPENSSL
|
||||
|| tls_ != nullptr || tls_local_ != nullptr
|
||||
#endif
|
||||
;
|
||||
}
|
||||
|
||||
/// @brief Underlying socket descriptor, or -1 for an inbox tape / TLS edge.
|
||||
/// @details Prefer framed `send` / `recv` on TLS channels. Raw descriptor
|
||||
/// I/O bypasses TLS and is rejected when the channel is encrypted.
|
||||
HEDLEY_NO_THROW
|
||||
int native_handle() noexcept
|
||||
{
|
||||
if (tcp_)
|
||||
return tcp_->native_handle();
|
||||
if (local_)
|
||||
return local_->native_handle();
|
||||
#if DPF_HAS_OPENSSL
|
||||
if (tls_ || tls_local_)
|
||||
return -1;
|
||||
#endif
|
||||
return -1;
|
||||
}
|
||||
|
||||
/// @brief True when application bytes ride TLS 1.3 under the frame header.
|
||||
HEDLEY_NO_THROW
|
||||
bool encrypted() const noexcept
|
||||
{
|
||||
#if DPF_HAS_OPENSSL
|
||||
return tls_ != nullptr || tls_local_ != nullptr;
|
||||
#else
|
||||
return false;
|
||||
#endif
|
||||
}
|
||||
|
||||
/// @brief Bytes and frames since construction or the last `reset_tally`.
|
||||
HEDLEY_NO_THROW
|
||||
io_tally tally() const noexcept
|
||||
{
|
||||
return tally_;
|
||||
}
|
||||
|
||||
/// @brief Zero the byte and frame counters. Does not touch the socket.
|
||||
HEDLEY_NO_THROW
|
||||
void reset_tally() noexcept
|
||||
{
|
||||
tally_ = {};
|
||||
}
|
||||
|
||||
/// @name Framed messages
|
||||
/// @brief Each frame is a little-endian length, a tag, then the payload.
|
||||
/// `T` must be trivially copyable. A zero-length payload is a
|
||||
/// header only.
|
||||
/// @throws std::runtime_error if the socket closes or the tag/size disagree
|
||||
/// @throws std::invalid_argument if a frame exceeds 2^32-1 bytes, or
|
||||
/// `exchange` is called with `self_id == peer_id`
|
||||
/// @throws std::logic_error if the channel is closed
|
||||
/// @{
|
||||
|
||||
/// @brief Send one value.
|
||||
/// @tparam T trivially copyable payload
|
||||
/// @param tag the message tag
|
||||
/// @param value the payload
|
||||
template <typename T>
|
||||
void send(msg tag, const T & value)
|
||||
{
|
||||
static_assert(std::is_trivially_copyable_v<T>,
|
||||
"net::channel::send requires a trivially copyable type");
|
||||
write_frame(tag, &value, sizeof(T));
|
||||
}
|
||||
|
||||
/// @brief Receive one value.
|
||||
/// @tparam T trivially copyable payload
|
||||
/// @param tag the expected message tag
|
||||
/// @return the payload
|
||||
template <typename T>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
T recv(msg tag)
|
||||
{
|
||||
static_assert(std::is_trivially_copyable_v<T>,
|
||||
"net::channel::recv requires a trivially copyable type");
|
||||
T value{};
|
||||
read_frame(tag, &value, sizeof(T));
|
||||
return value;
|
||||
}
|
||||
|
||||
/// @brief Send `n` values.
|
||||
/// @param data the values
|
||||
/// @param n the value count
|
||||
/// @param tag the message tag
|
||||
template <typename T>
|
||||
void send_vec(const T * data, std::size_t n, msg tag = msg::ring_vector)
|
||||
{
|
||||
static_assert(std::is_trivially_copyable_v<T>,
|
||||
"net::channel::send_vec requires a trivially copyable type");
|
||||
write_frame(tag, data, n * sizeof(T));
|
||||
}
|
||||
|
||||
/// @brief Send a vector of values.
|
||||
/// @param v the values
|
||||
/// @param tag the message tag
|
||||
template <typename T>
|
||||
void send_vec(const std::vector<T> & v, msg tag = msg::ring_vector)
|
||||
{
|
||||
send_vec(v.data(), v.size(), tag);
|
||||
}
|
||||
|
||||
/// @brief Receive a homogeneous vector.
|
||||
/// @tparam T trivially copyable element
|
||||
/// @param tag the expected message tag
|
||||
/// @return the elements. Empty when the payload length is 0.
|
||||
template <typename T>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::vector<T> recv_vec(msg tag = msg::ring_vector)
|
||||
{
|
||||
static_assert(std::is_trivially_copyable_v<T>,
|
||||
"net::channel::recv_vec requires a trivially copyable type");
|
||||
auto bytes = read_frame_bytes(tag);
|
||||
if (bytes.size() % sizeof(T) != 0)
|
||||
throw std::runtime_error("net::channel::recv_vec size mismatch");
|
||||
std::vector<T> out(bytes.size() / sizeof(T));
|
||||
if (!out.empty())
|
||||
std::memcpy(out.data(), bytes.data(), bytes.size());
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief Send `n` bytes.
|
||||
/// @param tag the message tag
|
||||
/// @param data the bytes
|
||||
/// @param n the byte count
|
||||
void send_bytes(msg tag, const void * data, std::size_t n)
|
||||
{
|
||||
write_frame(tag, data, n);
|
||||
}
|
||||
|
||||
/// @brief Send a byte vector.
|
||||
/// @param tag the message tag
|
||||
/// @param bytes the bytes
|
||||
void send_bytes(msg tag, const std::vector<std::uint8_t> & bytes)
|
||||
{
|
||||
send_bytes(tag, bytes.data(), bytes.size());
|
||||
}
|
||||
|
||||
/// @brief Receive an untyped payload.
|
||||
/// @param tag the expected message tag
|
||||
/// @return the payload bytes
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::vector<std::uint8_t> recv_bytes(msg tag)
|
||||
{
|
||||
return read_frame_bytes(tag);
|
||||
}
|
||||
|
||||
/// @brief Exchange one value. The lower id sends first.
|
||||
/// @tparam T trivially copyable payload
|
||||
/// @param self_id this party's role, as an integer
|
||||
/// @param peer_id the peer's role, as an integer
|
||||
/// @param mine this party's value
|
||||
/// @param tag the message tag
|
||||
/// @return the peer's value
|
||||
template <typename T>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
T exchange(unsigned self_id, unsigned peer_id, const T & mine, msg tag = msg::delta)
|
||||
{
|
||||
static_assert(std::is_trivially_copyable_v<T>,
|
||||
"net::channel::exchange requires a trivially copyable type");
|
||||
++tally_.exchanges;
|
||||
if (self_id < peer_id)
|
||||
{
|
||||
send(tag, mine);
|
||||
return recv<T>(tag);
|
||||
}
|
||||
if (self_id > peer_id)
|
||||
{
|
||||
T theirs = recv<T>(tag);
|
||||
send(tag, mine);
|
||||
return theirs;
|
||||
}
|
||||
throw std::invalid_argument("net::channel::exchange with self");
|
||||
}
|
||||
|
||||
/// @brief Exchange a homogeneous vector. One barrier; lower id sends first.
|
||||
/// @tparam T trivially copyable element
|
||||
/// @param self_id this party's role, as an integer
|
||||
/// @param peer_id the peer's role, as an integer
|
||||
/// @param mine this party's values
|
||||
/// @param tag the message tag
|
||||
/// @return the peer's values (same length as `mine`)
|
||||
/// @throws std::runtime_error if the peer's vector length disagrees
|
||||
template <typename T>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
std::vector<T> exchange_vec(unsigned self_id, unsigned peer_id,
|
||||
const std::vector<T> & mine, msg tag = msg::ring_vector)
|
||||
{
|
||||
static_assert(std::is_trivially_copyable_v<T>,
|
||||
"net::channel::exchange_vec requires a trivially copyable type");
|
||||
if (mine.empty())
|
||||
return {};
|
||||
++tally_.exchanges;
|
||||
std::vector<T> theirs;
|
||||
if (self_id < peer_id)
|
||||
{
|
||||
send_vec(mine, tag);
|
||||
theirs = recv_vec<T>(tag);
|
||||
}
|
||||
else if (self_id > peer_id)
|
||||
{
|
||||
theirs = recv_vec<T>(tag);
|
||||
send_vec(mine, tag);
|
||||
}
|
||||
else
|
||||
throw std::invalid_argument("net::channel::exchange_vec with self");
|
||||
if (theirs.size() != mine.size())
|
||||
throw std::runtime_error("net::channel::exchange_vec size mismatch");
|
||||
return theirs;
|
||||
}
|
||||
|
||||
/// @}
|
||||
|
||||
private:
|
||||
struct inbox_buf
|
||||
{
|
||||
std::vector<std::uint8_t> bytes;
|
||||
std::size_t pos = 0;
|
||||
};
|
||||
|
||||
std::unique_ptr<tcp_socket> tcp_;
|
||||
std::unique_ptr<local_socket> local_;
|
||||
#if DPF_HAS_OPENSSL
|
||||
std::unique_ptr<tls_stream> tls_;
|
||||
std::unique_ptr<tls_local_stream> tls_local_;
|
||||
std::shared_ptr<tls_context> tls_ctx_;
|
||||
#endif
|
||||
std::unique_ptr<inbox_buf> inbox_;
|
||||
io_tally tally_{};
|
||||
|
||||
void write_all(const void * data, std::size_t n)
|
||||
{
|
||||
if (inbox_)
|
||||
throw std::logic_error("net::channel inbox is recv-only");
|
||||
asio::const_buffer buf(data, n);
|
||||
asio::error_code ec;
|
||||
if (tcp_)
|
||||
asio::write(*tcp_, buf, ec);
|
||||
else if (local_)
|
||||
asio::write(*local_, buf, ec);
|
||||
#if DPF_HAS_OPENSSL
|
||||
else if (tls_)
|
||||
asio::write(*tls_, buf, ec);
|
||||
else if (tls_local_)
|
||||
asio::write(*tls_local_, buf, ec);
|
||||
#endif
|
||||
else
|
||||
throw std::logic_error("net::channel is closed");
|
||||
if (ec)
|
||||
throw std::runtime_error("net::channel write: " + ec.message());
|
||||
tally_.bytes_sent += n;
|
||||
}
|
||||
|
||||
void read_all(void * data, std::size_t n)
|
||||
{
|
||||
if (inbox_)
|
||||
{
|
||||
if (inbox_->pos + n > inbox_->bytes.size())
|
||||
throw std::runtime_error("net::channel inbox underrun");
|
||||
if (n != 0)
|
||||
std::memcpy(data, inbox_->bytes.data() + inbox_->pos, n);
|
||||
inbox_->pos += n;
|
||||
tally_.bytes_recv += n;
|
||||
return;
|
||||
}
|
||||
asio::mutable_buffer buf(data, n);
|
||||
asio::error_code ec;
|
||||
if (tcp_)
|
||||
asio::read(*tcp_, buf, ec);
|
||||
else if (local_)
|
||||
asio::read(*local_, buf, ec);
|
||||
#if DPF_HAS_OPENSSL
|
||||
else if (tls_)
|
||||
asio::read(*tls_, buf, ec);
|
||||
else if (tls_local_)
|
||||
asio::read(*tls_local_, buf, ec);
|
||||
#endif
|
||||
else
|
||||
throw std::logic_error("net::channel is closed");
|
||||
if (ec)
|
||||
throw std::runtime_error("net::channel read: " + ec.message());
|
||||
tally_.bytes_recv += n;
|
||||
}
|
||||
|
||||
void write_frame(msg tag, const void * payload, std::size_t n)
|
||||
{
|
||||
if (n > 0xffffffffu)
|
||||
throw std::invalid_argument("net::channel frame too large");
|
||||
std::uint32_t len = static_cast<std::uint32_t>(n);
|
||||
std::uint16_t t = to_u16(tag);
|
||||
std::array<std::uint8_t, 6> hdr{};
|
||||
std::memcpy(hdr.data(), &len, 4);
|
||||
std::memcpy(hdr.data() + 4, &t, 2);
|
||||
write_all(hdr.data(), hdr.size());
|
||||
if (n != 0)
|
||||
write_all(payload, n);
|
||||
tally_.payload_sent += n;
|
||||
++tally_.frames_sent;
|
||||
}
|
||||
|
||||
std::vector<std::uint8_t> read_frame_bytes(msg expected)
|
||||
{
|
||||
std::array<std::uint8_t, 6> hdr{};
|
||||
read_all(hdr.data(), hdr.size());
|
||||
std::uint32_t len = 0;
|
||||
std::uint16_t t = 0;
|
||||
std::memcpy(&len, hdr.data(), 4);
|
||||
std::memcpy(&t, hdr.data() + 4, 2);
|
||||
if (t != to_u16(expected))
|
||||
throw std::runtime_error("net::channel unexpected message tag");
|
||||
std::vector<std::uint8_t> body(len);
|
||||
if (len != 0)
|
||||
read_all(body.data(), body.size());
|
||||
tally_.payload_recv += len;
|
||||
++tally_.frames_recv;
|
||||
return body;
|
||||
}
|
||||
|
||||
void read_frame(msg expected, void * dest, std::size_t expect_n)
|
||||
{
|
||||
auto body = read_frame_bytes(expected);
|
||||
if (body.size() != expect_n)
|
||||
throw std::runtime_error("net::channel frame size mismatch");
|
||||
if (expect_n != 0)
|
||||
std::memcpy(dest, body.data(), expect_n);
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace net
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_NET_CHANNEL_HPP__
|
||||
218
include/dpf/net/client_link.hpp
Normal file
218
include/dpf/net/client_link.hpp
Normal file
|
|
@ -0,0 +1,218 @@
|
|||
/// @file dpf/net/client_link.hpp
|
||||
/// @brief Client-to-server links: TLS 1.3, the server always verified.
|
||||
/// @details A client that supplies inputs (for example, one share to each
|
||||
/// party) connects with `connect_server`. It verifies the server
|
||||
/// unless `client_security::verify` is off: a pinned server key, a CA
|
||||
/// chain for the host name, or, when neither is configured, the
|
||||
/// built-in development certificate, which a `client_listener` with no
|
||||
/// certificate or identity presents. The development certificate's
|
||||
/// private key is public, so that pairing works out of the box and is
|
||||
/// logged as providing no security. A server may also check client
|
||||
/// keys (`server_security::client_pins`). Both ends get an ordinary
|
||||
/// `async_stream_array` of `lanes` lanes and the link's
|
||||
/// `link_security`.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_NET_CLIENT_LINK_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_NET_CLIENT_LINK_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <memory>
|
||||
#include <stdexcept>
|
||||
#include <string>
|
||||
|
||||
#include "dpf/net/asio_ns.hpp"
|
||||
|
||||
#include "dpf/log.hpp"
|
||||
#include "dpf/net/async_stream_array.hpp"
|
||||
#include "dpf/net/connect.hpp"
|
||||
#include "dpf/net/link_log.hpp"
|
||||
#include "dpf/net/policy.hpp"
|
||||
#include "dpf/net/security.hpp"
|
||||
#include "dpf/net/socket_tune.hpp"
|
||||
#include "dpf/net/tls.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace net
|
||||
{
|
||||
|
||||
/// @brief One client link and how it was secured.
|
||||
struct client_connection
|
||||
{
|
||||
#if DPF_HAS_OPENSSL
|
||||
std::shared_ptr<tls_context> context;
|
||||
#endif
|
||||
std::unique_ptr<async_stream_array> link;
|
||||
link_security security;
|
||||
};
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
inline constexpr std::uint32_t client_magic = 0x4c435044u; // 'DPCL'
|
||||
|
||||
#if DPF_HAS_OPENSSL
|
||||
/// @brief Both ends send `{magic, lanes}` inside TLS and must agree.
|
||||
inline void client_hello(asio::io_context & io, tls_stream & s, std::size_t lanes,
|
||||
bool server, std::chrono::milliseconds budget, const std::string & who)
|
||||
{
|
||||
std::uint8_t mine[8];
|
||||
std::uint8_t theirs[8];
|
||||
put_u32(mine, client_magic);
|
||||
put_u32(mine + 4, static_cast<std::uint32_t>(lanes));
|
||||
if (server)
|
||||
{
|
||||
tls_read(io, s, theirs, sizeof(theirs), budget, who);
|
||||
tls_write(io, s, mine, sizeof(mine), budget, who);
|
||||
}
|
||||
else
|
||||
{
|
||||
tls_write(io, s, mine, sizeof(mine), budget, who);
|
||||
tls_read(io, s, theirs, sizeof(theirs), budget, who);
|
||||
}
|
||||
if (get_u32(theirs) != client_magic)
|
||||
throw std::runtime_error(who + ": the peer is not a libdpf client link");
|
||||
if (get_u32(theirs + 4) != lanes)
|
||||
throw std::runtime_error(who + ": lanes " + std::to_string(get_u32(theirs + 4))
|
||||
+ " vs " + std::to_string(lanes) + " (peer vs this side)");
|
||||
}
|
||||
#endif
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/// @brief Connect to the server at `host:port` and verify it.
|
||||
inline client_connection connect_server(asio::io_context & io, const std::string & host,
|
||||
unsigned short port, const client_security & sec, std::size_t lanes = 1,
|
||||
const wire_policy & pol = {}, const deadlines & lim = {})
|
||||
{
|
||||
#if DPF_HAS_OPENSSL
|
||||
const std::string where = host + ":" + std::to_string(port);
|
||||
client_connection out;
|
||||
out.context = make_client_tls_context(sec);
|
||||
asio::ip::tcp::socket sock(io);
|
||||
connect_until(sock, host, port, lim.connect);
|
||||
tune_tcp(sock, pol.socket);
|
||||
const int fd = sock.native_handle();
|
||||
tls_stream s(std::move(sock), *out.context);
|
||||
if (sec.verify && !sec.ca_file.empty())
|
||||
tls_expect_host(s, sec.server_name.empty() ? host : sec.server_name);
|
||||
tls_handshake(io, s, false, lim.handshake, "client: TLS handshake with " + where);
|
||||
out.security = tls_describe(s);
|
||||
try
|
||||
{
|
||||
check_server(out.security, s, sec, where);
|
||||
}
|
||||
catch (...)
|
||||
{
|
||||
std::error_code e;
|
||||
s.lowest_layer().close(e);
|
||||
throw;
|
||||
}
|
||||
if (!sec.verify)
|
||||
DPF_LOG(error, "client.verify_off").kv("server", where)
|
||||
.kv("detail", "client_verify=off: any server certificate is accepted, so "
|
||||
"this connection is encrypted but the server is not authenticated");
|
||||
else if (out.security.peer_auth == "development"
|
||||
&& log::first_time("client.development." + where))
|
||||
DPF_LOG(warning, "client.development_certificate").kv("server", where)
|
||||
.kv("detail", "the server presented the built-in development certificate, "
|
||||
"whose private key is public: this connection is encrypted but the "
|
||||
"server is not authenticated (pin its key or configure client_ca)");
|
||||
detail::client_hello(io, s, lanes, false, lim.handshake, "client link to " + where);
|
||||
out.link = std::make_unique<async_tls_mux_stream_array>(io, std::move(s), 1, 0, lanes,
|
||||
pol);
|
||||
log_link_up("connect", "server", transport::mux, lanes, 0, 0, fd, pol.socket,
|
||||
&out.security);
|
||||
return out;
|
||||
#else
|
||||
(void)io;
|
||||
(void)host;
|
||||
(void)port;
|
||||
(void)sec;
|
||||
(void)lanes;
|
||||
(void)pol;
|
||||
(void)lim;
|
||||
throw std::logic_error("connect_server: built without OpenSSL");
|
||||
#endif
|
||||
}
|
||||
|
||||
/// @brief Server side: accept clients, each on its own TLS link.
|
||||
class client_listener
|
||||
{
|
||||
public:
|
||||
client_listener(asio::io_context & io, server_security sec, std::size_t lanes = 1,
|
||||
wire_policy pol = {}, deadlines lim = {})
|
||||
: io_(&io), sec_(std::move(sec)), lanes_(lanes), pol_(pol), lim_(lim)
|
||||
{
|
||||
#if DPF_HAS_OPENSSL
|
||||
ctx_ = make_server_tls_context(sec_, development_);
|
||||
if (development_ && log::first_time("server.development"))
|
||||
DPF_LOG(warning, "server.development_certificate")
|
||||
.kv("detail", "presenting the built-in development certificate, whose "
|
||||
"private key is public: clients cannot tell this server from any "
|
||||
"other (set server_cert/server_key or server_identity)");
|
||||
#else
|
||||
throw std::logic_error("client_listener: built without OpenSSL");
|
||||
#endif
|
||||
}
|
||||
|
||||
/// @brief Bind `port` (0 = ephemeral) and return it.
|
||||
unsigned short listen(unsigned short port = 0)
|
||||
{
|
||||
if (!acceptor_)
|
||||
{
|
||||
acceptor_ = std::make_unique<asio::ip::tcp::acceptor>(*io_);
|
||||
open_listener(*acceptor_, port);
|
||||
log_listen(acceptor_->local_endpoint().port(), false, true);
|
||||
}
|
||||
return acceptor_->local_endpoint().port();
|
||||
}
|
||||
|
||||
bool development() const noexcept { return development_; }
|
||||
|
||||
/// @brief Wait (up to the accept deadline) for one client.
|
||||
client_connection accept()
|
||||
{
|
||||
listen(0);
|
||||
client_connection out;
|
||||
#if DPF_HAS_OPENSSL
|
||||
out.context = ctx_;
|
||||
asio::ip::tcp::socket sock(*io_);
|
||||
accept_until(*acceptor_, sock, lim_.accept);
|
||||
tune_tcp(sock, pol_.socket);
|
||||
const int fd = sock.native_handle();
|
||||
tls_stream s(std::move(sock), *ctx_);
|
||||
tls_handshake(*io_, s, true, lim_.handshake, "server: TLS handshake with a client");
|
||||
out.security = tls_describe(s);
|
||||
check_client(out.security, sec_);
|
||||
detail::client_hello(*io_, s, lanes_, true, lim_.handshake, "client link");
|
||||
out.link = std::make_unique<async_tls_mux_stream_array>(*io_, std::move(s), 0, 1,
|
||||
lanes_, pol_);
|
||||
log_link_up("accept", "client", transport::mux, lanes_, 0, 0, fd, pol_.socket,
|
||||
&out.security);
|
||||
if (!sec_.client_pins.empty() && out.security.peer_auth == "none")
|
||||
DPF_LOG(warning, "server.client_unauthenticated")
|
||||
.kv("client_key", out.security.peer_key ? out.security.peer_key->base64()
|
||||
: std::string("none"))
|
||||
.kv("detail", "the client presented no pinned key");
|
||||
#endif
|
||||
return out;
|
||||
}
|
||||
|
||||
private:
|
||||
asio::io_context * io_ = nullptr;
|
||||
server_security sec_;
|
||||
std::size_t lanes_ = 1;
|
||||
wire_policy pol_{};
|
||||
deadlines lim_{};
|
||||
bool development_ = false;
|
||||
#if DPF_HAS_OPENSSL
|
||||
std::shared_ptr<tls_context> ctx_;
|
||||
#endif
|
||||
std::unique_ptr<asio::ip::tcp::acceptor> acceptor_;
|
||||
};
|
||||
|
||||
} // namespace net
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_NET_CLIENT_LINK_HPP__
|
||||
58
include/dpf/net/comm_hook.hpp
Normal file
58
include/dpf/net/comm_hook.hpp
Normal file
|
|
@ -0,0 +1,58 @@
|
|||
/// @file dpf/net/comm_hook.hpp
|
||||
/// @brief Replaceable transport under trio helpers.
|
||||
/// @details Protocols call `trio::exchange_with`, `send_to` / `recv_from`,
|
||||
/// `deal` / `accept_deal`, and `trio::batch`. Those helpers go through
|
||||
/// a `comm_hook` when one is installed. A null hook keeps the current
|
||||
/// framed mesh. `channel::send` / `recv` stay the raw bypass.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_NET_COMM_HOOK_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_NET_COMM_HOOK_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <memory>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/net/channel.hpp"
|
||||
#include "dpf/net/round_sink.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace net
|
||||
{
|
||||
|
||||
class trio;
|
||||
|
||||
/// @brief Byte-level transport for trio helpers and RoundSink factories.
|
||||
/// @details Overrides replace the mesh without changing protocol call sites.
|
||||
/// Peer ids are `0..2` (`role` as unsigned). Default implementations
|
||||
/// in `mesh_comm_hook` match today's `channel` framing.
|
||||
class comm_hook
|
||||
{
|
||||
public:
|
||||
virtual ~comm_hook() = default;
|
||||
|
||||
/// @brief One-way send on the link to `peer_id`.
|
||||
virtual void send_bytes(trio & net, unsigned peer_id, msg tag,
|
||||
const void * data, std::size_t n) = 0;
|
||||
|
||||
/// @brief One-way receive from `peer_id`.
|
||||
virtual std::vector<std::uint8_t> recv_bytes(trio & net, unsigned peer_id,
|
||||
msg tag) = 0;
|
||||
|
||||
/// @brief Duplex exchange. Lower role id sends first.
|
||||
virtual std::vector<std::uint8_t> exchange_bytes(trio & net,
|
||||
unsigned peer_id, msg tag, const void * data, std::size_t n) = 0;
|
||||
|
||||
/// @brief Homogeneous vector exchange (one frame each way, same length).
|
||||
virtual std::vector<std::uint8_t> exchange_vec_bytes(trio & net,
|
||||
unsigned peer_id, msg tag, const void * data, std::size_t nbytes) = 0;
|
||||
|
||||
/// @brief Round-batched peer sink for `count` instances.
|
||||
virtual std::unique_ptr<RoundSink> batch(trio & net, std::size_t count,
|
||||
std::vector<std::size_t> slot_bytes) = 0;
|
||||
};
|
||||
|
||||
} // namespace net
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_NET_COMM_HOOK_HPP__
|
||||
263
include/dpf/net/connect.hpp
Normal file
263
include/dpf/net/connect.hpp
Normal file
|
|
@ -0,0 +1,263 @@
|
|||
/// @file dpf/net/connect.hpp
|
||||
/// @brief Deadline-bounded connect, accept, and handshake I/O for setup.
|
||||
/// @details Setup is blocking by design (it happens before protocol traffic),
|
||||
/// but every step has a wall-clock bound. `connect_until` retries
|
||||
/// refused connections until the deadline, so processes may start in
|
||||
/// any order.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_NET_CONNECT_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_NET_CONNECT_HPP__
|
||||
|
||||
#include <cerrno>
|
||||
#include <chrono>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <stdexcept>
|
||||
#include <string>
|
||||
#include <system_error>
|
||||
#include <thread>
|
||||
|
||||
#include <fcntl.h>
|
||||
#include <poll.h>
|
||||
#include <sys/socket.h>
|
||||
#include <sys/types.h>
|
||||
#include <unistd.h>
|
||||
|
||||
#include "dpf/net/asio_ns.hpp"
|
||||
|
||||
#include "dpf/log.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace net
|
||||
{
|
||||
|
||||
using setup_clock = std::chrono::steady_clock;
|
||||
|
||||
namespace detail
|
||||
{
|
||||
|
||||
inline int remaining_ms(setup_clock::time_point deadline)
|
||||
{
|
||||
const auto left = std::chrono::duration_cast<std::chrono::milliseconds>(
|
||||
deadline - setup_clock::now());
|
||||
if (left.count() <= 0)
|
||||
return 0;
|
||||
if (left.count() > 0x7fffffff)
|
||||
return 0x7fffffff;
|
||||
return static_cast<int>(left.count());
|
||||
}
|
||||
|
||||
/// @brief Wait until `fd` is ready for `events` or throw at `deadline`.
|
||||
inline void wait_fd(int fd, short events, setup_clock::time_point deadline,
|
||||
const std::string & what)
|
||||
{
|
||||
for (;;)
|
||||
{
|
||||
const int ms = remaining_ms(deadline);
|
||||
if (ms == 0)
|
||||
throw std::system_error(std::make_error_code(std::errc::timed_out),
|
||||
what + ": timed out");
|
||||
pollfd pfd{};
|
||||
pfd.fd = fd;
|
||||
pfd.events = events;
|
||||
const int rc = ::poll(&pfd, 1, ms);
|
||||
if (rc > 0)
|
||||
{
|
||||
if ((pfd.revents & (POLLERR | POLLNVAL)) != 0
|
||||
&& (pfd.revents & events) == 0)
|
||||
throw std::system_error(
|
||||
std::make_error_code(std::errc::connection_reset), what);
|
||||
return;
|
||||
}
|
||||
if (rc < 0 && errno != EINTR)
|
||||
throw std::system_error(errno, std::generic_category(), what);
|
||||
}
|
||||
}
|
||||
|
||||
inline void send_all_until(int fd, const void * data, std::size_t n,
|
||||
setup_clock::time_point deadline, const std::string & what)
|
||||
{
|
||||
const auto * p = static_cast<const std::uint8_t *>(data);
|
||||
std::size_t done = 0;
|
||||
while (done < n)
|
||||
{
|
||||
const ssize_t r = ::send(fd, p + done, n - done,
|
||||
MSG_DONTWAIT | MSG_NOSIGNAL);
|
||||
if (r > 0)
|
||||
{
|
||||
done += static_cast<std::size_t>(r);
|
||||
continue;
|
||||
}
|
||||
if (r < 0 && (errno == EAGAIN || errno == EWOULDBLOCK || errno == EINTR))
|
||||
{
|
||||
wait_fd(fd, POLLOUT, deadline, what);
|
||||
continue;
|
||||
}
|
||||
throw std::system_error(errno, std::generic_category(), what);
|
||||
}
|
||||
}
|
||||
|
||||
inline void recv_all_until(int fd, void * data, std::size_t n,
|
||||
setup_clock::time_point deadline, const std::string & what)
|
||||
{
|
||||
auto * p = static_cast<std::uint8_t *>(data);
|
||||
std::size_t done = 0;
|
||||
while (done < n)
|
||||
{
|
||||
const ssize_t r = ::recv(fd, p + done, n - done, MSG_DONTWAIT);
|
||||
if (r > 0)
|
||||
{
|
||||
done += static_cast<std::size_t>(r);
|
||||
continue;
|
||||
}
|
||||
if (r == 0)
|
||||
throw std::system_error(std::make_error_code(
|
||||
std::errc::connection_reset), what + ": peer closed");
|
||||
if (errno == EAGAIN || errno == EWOULDBLOCK || errno == EINTR)
|
||||
{
|
||||
wait_fd(fd, POLLIN, deadline, what);
|
||||
continue;
|
||||
}
|
||||
throw std::system_error(errno, std::generic_category(), what);
|
||||
}
|
||||
}
|
||||
|
||||
inline void put_u32(std::uint8_t * dst, std::uint32_t v) noexcept
|
||||
{
|
||||
dst[0] = static_cast<std::uint8_t>(v & 0xffu);
|
||||
dst[1] = static_cast<std::uint8_t>((v >> 8) & 0xffu);
|
||||
dst[2] = static_cast<std::uint8_t>((v >> 16) & 0xffu);
|
||||
dst[3] = static_cast<std::uint8_t>((v >> 24) & 0xffu);
|
||||
}
|
||||
|
||||
inline std::uint32_t get_u32(const std::uint8_t * src) noexcept
|
||||
{
|
||||
return static_cast<std::uint32_t>(src[0]) | (static_cast<std::uint32_t>(src[1]) << 8)
|
||||
| (static_cast<std::uint32_t>(src[2]) << 16)
|
||||
| (static_cast<std::uint32_t>(src[3]) << 24);
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/// @brief Connect `sock` to `host:port`, retrying refusals until `budget`.
|
||||
inline void connect_until(asio::ip::tcp::socket & sock, const std::string & host,
|
||||
unsigned short port, std::chrono::milliseconds budget)
|
||||
{
|
||||
const auto started = setup_clock::now();
|
||||
const auto deadline = started + budget;
|
||||
const std::string what = "connect " + host + ":" + std::to_string(port);
|
||||
asio::ip::tcp::resolver res(sock.get_executor());
|
||||
std::error_code rec;
|
||||
auto eps = res.resolve(host, std::to_string(port), rec);
|
||||
if (rec)
|
||||
throw std::system_error(rec, what + ": resolve");
|
||||
std::error_code last = std::make_error_code(std::errc::timed_out);
|
||||
std::size_t attempts = 0;
|
||||
for (;;)
|
||||
{
|
||||
for (const auto & entry : eps)
|
||||
{
|
||||
std::error_code ec;
|
||||
if (sock.is_open())
|
||||
sock.close(ec);
|
||||
sock.open(entry.endpoint().protocol(), ec);
|
||||
if (ec)
|
||||
{
|
||||
last = ec;
|
||||
continue;
|
||||
}
|
||||
const int fd = sock.native_handle();
|
||||
const int fl = ::fcntl(fd, F_GETFL, 0);
|
||||
(void)::fcntl(fd, F_SETFL, fl | O_NONBLOCK);
|
||||
++attempts;
|
||||
int rc = ::connect(fd, entry.endpoint().data(),
|
||||
static_cast<socklen_t>(entry.endpoint().size()));
|
||||
int err = rc == 0 ? 0 : errno;
|
||||
if (rc != 0 && err == EINPROGRESS)
|
||||
{
|
||||
pollfd pfd{};
|
||||
pfd.fd = fd;
|
||||
pfd.events = POLLOUT;
|
||||
const int prc = ::poll(&pfd, 1, detail::remaining_ms(deadline));
|
||||
if (prc <= 0)
|
||||
err = ETIMEDOUT;
|
||||
else
|
||||
{
|
||||
socklen_t len = sizeof(err);
|
||||
err = 0;
|
||||
::getsockopt(fd, SOL_SOCKET, SO_ERROR, &err, &len);
|
||||
}
|
||||
}
|
||||
if (err == 0)
|
||||
{
|
||||
(void)::fcntl(fd, F_SETFL, fl);
|
||||
DPF_LOG(debug, "connect").kv("target", host + ":" + std::to_string(port))
|
||||
.kv("resolved", entry.endpoint().address().to_string())
|
||||
.kv("attempts", attempts)
|
||||
.kv("elapsed_ms", std::chrono::duration<double, std::milli>(
|
||||
setup_clock::now() - started).count());
|
||||
return;
|
||||
}
|
||||
last = std::error_code(err, std::generic_category());
|
||||
sock.close(ec);
|
||||
}
|
||||
if (setup_clock::now() >= deadline)
|
||||
{
|
||||
DPF_LOG(error, "connect.failed").kv("target", host + ":" + std::to_string(port))
|
||||
.kv("attempts", attempts).kv("budget_ms", budget.count())
|
||||
.kv("last_error", last.message());
|
||||
throw std::system_error(last, what + " failed after "
|
||||
+ std::to_string(budget.count()) + " ms");
|
||||
}
|
||||
std::this_thread::sleep_for(std::chrono::milliseconds(20));
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Accept one connection on `acc` within `budget`.
|
||||
inline void accept_until(asio::ip::tcp::acceptor & acc,
|
||||
asio::ip::tcp::socket & sock, std::chrono::milliseconds budget)
|
||||
{
|
||||
const auto deadline = setup_clock::now() + budget;
|
||||
const std::string what = "accept on port "
|
||||
+ std::to_string(acc.local_endpoint().port());
|
||||
detail::wait_fd(acc.native_handle(), POLLIN, deadline, what);
|
||||
acc.accept(sock);
|
||||
}
|
||||
|
||||
/// @brief Open, bind, and listen on `port` (0 = ephemeral), reusing the address.
|
||||
inline void open_listener(asio::ip::tcp::acceptor & acc, unsigned short port)
|
||||
{
|
||||
const asio::ip::tcp::endpoint ep(asio::ip::tcp::v4(), port);
|
||||
acc.open(ep.protocol());
|
||||
acc.set_option(asio::socket_base::reuse_address(true));
|
||||
acc.bind(ep);
|
||||
acc.listen();
|
||||
}
|
||||
|
||||
/// @brief Exchange one `u32` each way within `budget` (setup handshake).
|
||||
inline std::uint32_t exchange_u32(int fd, std::uint32_t mine,
|
||||
std::chrono::milliseconds budget, const std::string & what)
|
||||
{
|
||||
const auto deadline = setup_clock::now() + budget;
|
||||
std::uint8_t out[4];
|
||||
detail::put_u32(out, mine);
|
||||
detail::send_all_until(fd, out, 4, deadline, what);
|
||||
std::uint8_t in[4];
|
||||
detail::recv_all_until(fd, in, 4, deadline, what);
|
||||
return detail::get_u32(in);
|
||||
}
|
||||
|
||||
/// @brief Send then receive a fixed-size record within `budget`.
|
||||
inline void exchange_record(int fd, const void * mine, void * theirs,
|
||||
std::size_t n, std::chrono::milliseconds budget, const std::string & what)
|
||||
{
|
||||
const auto deadline = setup_clock::now() + budget;
|
||||
detail::send_all_until(fd, mine, n, deadline, what);
|
||||
detail::recv_all_until(fd, theirs, n, deadline, what);
|
||||
}
|
||||
|
||||
} // namespace net
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_NET_CONNECT_HPP__
|
||||
221
include/dpf/net/dealer_cursor.hpp
Normal file
221
include/dpf/net/dealer_cursor.hpp
Normal file
|
|
@ -0,0 +1,221 @@
|
|||
/// @file dpf/net/dealer_cursor.hpp
|
||||
/// @brief Per-(round, index) correction words for a batch session.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_NET_DEALER_CURSOR_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_NET_DEALER_CURSOR_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <stdexcept>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/net/stream_array.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace net
|
||||
{
|
||||
|
||||
/// @brief Empty correction / blind marker. Protocols that need no dealer
|
||||
/// material specialize on this type and skip the cursor.
|
||||
struct empty_pad
|
||||
{
|
||||
};
|
||||
|
||||
/// @brief Seekable table of dealer corrections, one lane per round.
|
||||
/// @details Size is fixed at construction: `count` instances and one slot
|
||||
/// width per round. Live dealers and inbox tapes both fill this
|
||||
/// table before the online session runs, or fill it lazily through
|
||||
/// `put`. Sessions read with `at(round, index)`.
|
||||
class dealer_cursor
|
||||
{
|
||||
public:
|
||||
dealer_cursor() = default;
|
||||
|
||||
dealer_cursor(std::size_t count, std::vector<std::size_t> slot_bytes)
|
||||
: count_(count),
|
||||
slot_bytes_(std::move(slot_bytes)),
|
||||
offsets_(slot_bytes_.size() + 1, 0)
|
||||
{
|
||||
for (std::size_t r = 0; r < slot_bytes_.size(); ++r)
|
||||
offsets_[r + 1] = offsets_[r] + count_ * slot_bytes_[r];
|
||||
store_.assign(offsets_.back(), 0);
|
||||
ready_.assign(slot_bytes_.size() * count_, false);
|
||||
}
|
||||
|
||||
std::size_t count() const noexcept { return count_; }
|
||||
std::size_t rounds() const noexcept { return slot_bytes_.size(); }
|
||||
|
||||
std::size_t slot_bytes(std::uint16_t round) const
|
||||
{
|
||||
if (round >= slot_bytes_.size())
|
||||
throw std::out_of_range("dealer_cursor round");
|
||||
return slot_bytes_[round];
|
||||
}
|
||||
|
||||
void put(std::uint16_t round, std::size_t index, const std::uint8_t * bytes,
|
||||
std::size_t n)
|
||||
{
|
||||
check(round, index, n);
|
||||
std::memcpy(ptr(round, index), bytes, n);
|
||||
ready_[ready_index(round, index)] = true;
|
||||
}
|
||||
|
||||
/// @brief Fill every index of `round` from a contiguous dealer tape.
|
||||
void put_round(std::uint16_t round, const std::uint8_t * bytes,
|
||||
std::size_t nbytes)
|
||||
{
|
||||
if (round >= slot_bytes_.size())
|
||||
throw std::out_of_range("dealer_cursor round");
|
||||
const std::size_t need = count_ * slot_bytes_[round];
|
||||
if (nbytes != need)
|
||||
throw std::invalid_argument("dealer_cursor put_round size");
|
||||
if (need != 0)
|
||||
std::memcpy(store_.data() + offsets_[round], bytes, need);
|
||||
for (std::size_t i = 0; i < count_; ++i)
|
||||
ready_[ready_index(round, i)] = true;
|
||||
}
|
||||
|
||||
bool ready(std::uint16_t round, std::size_t index) const
|
||||
{
|
||||
if (round >= slot_bytes_.size() || index >= count_)
|
||||
return false;
|
||||
if (slot_bytes_[round] == 0)
|
||||
return true;
|
||||
return ready_[ready_index(round, index)];
|
||||
}
|
||||
|
||||
void at(std::uint16_t round, std::size_t index, std::uint8_t * out,
|
||||
std::size_t n) const
|
||||
{
|
||||
check(round, index, n);
|
||||
if (n == 0)
|
||||
return;
|
||||
if (!ready_[ready_index(round, index)])
|
||||
throw std::logic_error("dealer_cursor not ready");
|
||||
std::memcpy(out, ptr(round, index), n);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
T at(std::uint16_t round, std::size_t index) const
|
||||
{
|
||||
static_assert(std::is_trivially_copyable_v<T>,
|
||||
"dealer_cursor::at requires a trivially copyable type");
|
||||
if constexpr (std::is_same_v<T, empty_pad>)
|
||||
{
|
||||
(void)round;
|
||||
(void)index;
|
||||
return empty_pad{};
|
||||
}
|
||||
else
|
||||
{
|
||||
T out{};
|
||||
at(round, index, reinterpret_cast<std::uint8_t *>(&out), sizeof(T));
|
||||
return out;
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Load each round from consecutive stream indexes (one tape per round).
|
||||
static dealer_cursor from_streams(stream_array & src, std::size_t count,
|
||||
std::vector<std::size_t> slot_bytes, std::size_t first_stream = 0)
|
||||
{
|
||||
dealer_cursor out(count, std::move(slot_bytes));
|
||||
for (std::uint16_t r = 0; r < out.rounds(); ++r)
|
||||
{
|
||||
const std::size_t need = count * out.slot_bytes(r);
|
||||
if (need == 0)
|
||||
continue;
|
||||
std::vector<std::uint8_t> buf(need);
|
||||
src.read(first_stream + r, buf.data(), need);
|
||||
out.put_round(r, buf.data(), need);
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
private:
|
||||
void check(std::uint16_t round, std::size_t index, std::size_t n) const
|
||||
{
|
||||
if (round >= slot_bytes_.size() || index >= count_)
|
||||
throw std::out_of_range("dealer_cursor index");
|
||||
if (n != slot_bytes_[round])
|
||||
throw std::invalid_argument("dealer_cursor size");
|
||||
}
|
||||
|
||||
std::size_t ready_index(std::uint16_t round, std::size_t index) const noexcept
|
||||
{
|
||||
return static_cast<std::size_t>(round) * count_ + index;
|
||||
}
|
||||
|
||||
std::uint8_t * ptr(std::uint16_t round, std::size_t index) noexcept
|
||||
{
|
||||
return store_.data() + offsets_[round] + index * slot_bytes_[round];
|
||||
}
|
||||
|
||||
const std::uint8_t * ptr(std::uint16_t round, std::size_t index) const noexcept
|
||||
{
|
||||
return store_.data() + offsets_[round] + index * slot_bytes_[round];
|
||||
}
|
||||
|
||||
std::size_t count_ = 0;
|
||||
std::vector<std::size_t> slot_bytes_;
|
||||
std::vector<std::size_t> offsets_;
|
||||
std::vector<std::uint8_t> store_;
|
||||
std::vector<bool> ready_;
|
||||
};
|
||||
|
||||
/// @brief Prefetch a full dealer table from `stream_array` tapes (see `from_streams`).
|
||||
class stream_dealer_cursor
|
||||
{
|
||||
public:
|
||||
stream_dealer_cursor(stream_array & src, std::size_t count,
|
||||
std::vector<std::size_t> slot_bytes, std::size_t first_stream = 0)
|
||||
: inner_(dealer_cursor::from_streams(src, count, std::move(slot_bytes),
|
||||
first_stream))
|
||||
{
|
||||
}
|
||||
|
||||
std::size_t count() const noexcept { return inner_.count(); }
|
||||
std::size_t rounds() const noexcept { return inner_.rounds(); }
|
||||
|
||||
std::size_t slot_bytes(std::uint16_t round) const
|
||||
{
|
||||
return inner_.slot_bytes(round);
|
||||
}
|
||||
|
||||
bool ready(std::uint16_t round, std::size_t index) const
|
||||
{
|
||||
return inner_.ready(round, index);
|
||||
}
|
||||
|
||||
void at(std::uint16_t round, std::size_t index, std::uint8_t * out,
|
||||
std::size_t n) const
|
||||
{
|
||||
inner_.at(round, index, out, n);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
T at(std::uint16_t round, std::size_t index) const
|
||||
{
|
||||
return inner_.at<T>(round, index);
|
||||
}
|
||||
|
||||
const dealer_cursor & cursor() const noexcept { return inner_; }
|
||||
dealer_cursor & cursor() noexcept { return inner_; }
|
||||
|
||||
private:
|
||||
dealer_cursor inner_;
|
||||
};
|
||||
|
||||
/// @brief One-round table: `nbytes` must equal `count * slot_bytes[0]`.
|
||||
inline dealer_cursor dealer_cursor_from_stream(stream_array & a,
|
||||
std::size_t stream_index, std::size_t count, std::size_t slot_nbytes)
|
||||
{
|
||||
return dealer_cursor::from_streams(a, count, {slot_nbytes}, stream_index);
|
||||
}
|
||||
|
||||
} // namespace net
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_NET_DEALER_CURSOR_HPP__
|
||||
246
include/dpf/net/edge_mesh.hpp
Normal file
246
include/dpf/net/edge_mesh.hpp
Normal file
|
|
@ -0,0 +1,246 @@
|
|||
/// @file dpf/net/edge_mesh.hpp
|
||||
/// @brief N-edge RoundSink mesh for star / dealer / 4PC topologies.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_NET_EDGE_MESH_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_NET_EDGE_MESH_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <memory>
|
||||
#include <stdexcept>
|
||||
#include <string>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/net/memory_sink.hpp"
|
||||
#include "dpf/net/round_sink.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace net
|
||||
{
|
||||
|
||||
/// @brief Index of a duplex link in an `edge_mesh`.
|
||||
using edge_id = std::uint16_t;
|
||||
|
||||
/// @brief Named edges for the common 2PC / RSS / dealer trio (also mesh ids 0..2).
|
||||
/// @details For the dealer (party 2), `edge_dealer` is its link to party 0 and
|
||||
/// `edge_dealer_p1` its link to party 1.
|
||||
inline constexpr edge_id edge_peer = 0;
|
||||
inline constexpr edge_id edge_rss_next = 1;
|
||||
inline constexpr edge_id edge_dealer = 2;
|
||||
inline constexpr edge_id edge_dealer_p1 = 3;
|
||||
|
||||
inline std::string edge_name(edge_id e)
|
||||
{
|
||||
switch (e)
|
||||
{
|
||||
case edge_peer:
|
||||
return "peer";
|
||||
case edge_rss_next:
|
||||
return "rss_next";
|
||||
case edge_dealer:
|
||||
return "dealer";
|
||||
case edge_dealer_p1:
|
||||
return "dealer->p1";
|
||||
default:
|
||||
return "edge " + std::to_string(e);
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Collection of duplex RoundSinks keyed by `edge_id`.
|
||||
struct edge_mesh
|
||||
{
|
||||
std::vector<RoundSink *> sinks;
|
||||
|
||||
RoundSink & at(edge_id id) const
|
||||
{
|
||||
if (static_cast<std::size_t>(id) >= sinks.size() || sinks[id] == nullptr)
|
||||
throw std::logic_error("edge_mesh: edge not bound");
|
||||
return *sinks[id];
|
||||
}
|
||||
|
||||
bool has(edge_id id) const noexcept
|
||||
{
|
||||
return static_cast<std::size_t>(id) < sinks.size()
|
||||
&& sinks[id] != nullptr;
|
||||
}
|
||||
|
||||
std::size_t size() const noexcept { return sinks.size(); }
|
||||
|
||||
void flush_all()
|
||||
{
|
||||
for (auto * s : sinks)
|
||||
{
|
||||
if (s == nullptr)
|
||||
continue;
|
||||
s->flush();
|
||||
s->poll();
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief Sum of `progress()` over distinct bound sinks.
|
||||
std::uint64_t progress_total() const noexcept
|
||||
{
|
||||
std::uint64_t total = 0;
|
||||
for (std::size_t i = 0; i < sinks.size(); ++i)
|
||||
{
|
||||
const RoundSink * s = sinks[i];
|
||||
if (s == nullptr)
|
||||
continue;
|
||||
bool seen = false;
|
||||
for (std::size_t j = 0; j < i && !seen; ++j)
|
||||
seen = sinks[j] == s;
|
||||
if (!seen)
|
||||
total += s->progress();
|
||||
}
|
||||
return total;
|
||||
}
|
||||
|
||||
/// @brief Block until at least one bound sink makes I/O progress.
|
||||
/// @details Returns true if any sink's `wait_io()` ran a real event (an
|
||||
/// async sink slept in `epoll`); false if none did (memory sinks),
|
||||
/// so the caller can fall back to its spin guard.
|
||||
bool wait_io_all()
|
||||
{
|
||||
bool progressed = false;
|
||||
for (auto * s : sinks)
|
||||
{
|
||||
if (s == nullptr)
|
||||
continue;
|
||||
if (s->wait_io())
|
||||
progressed = true;
|
||||
}
|
||||
return progressed;
|
||||
}
|
||||
};
|
||||
|
||||
/// @brief In-process star: one client edge per server, matching server ends.
|
||||
struct memory_star
|
||||
{
|
||||
std::size_t servers = 0;
|
||||
std::vector<std::shared_ptr<memory_sink_hub>> hubs;
|
||||
std::vector<memory_sink> client; ///< client side of edge i
|
||||
std::vector<memory_sink> server; ///< server i side of edge i
|
||||
|
||||
/// @brief Client mesh: edges `[0, servers)`.
|
||||
edge_mesh client_mesh()
|
||||
{
|
||||
edge_mesh m;
|
||||
m.sinks.resize(client.size());
|
||||
for (std::size_t i = 0; i < client.size(); ++i)
|
||||
m.sinks[i] = &client[i];
|
||||
return m;
|
||||
}
|
||||
|
||||
/// @brief Server `i` mesh with a single live edge at `edge_id{i}` (sparse).
|
||||
/// Prefer `server_edge(i)` when the schedule uses `edge_id{0}` locally.
|
||||
edge_mesh server_mesh_at(std::size_t i)
|
||||
{
|
||||
if (i >= server.size())
|
||||
throw std::out_of_range("memory_star server");
|
||||
edge_mesh m;
|
||||
m.sinks.assign(server.size(), nullptr);
|
||||
m.sinks[i] = &server[i];
|
||||
return m;
|
||||
}
|
||||
|
||||
/// @brief Server `i` as a one-edge mesh (`edge_id` 0 → that duplex).
|
||||
edge_mesh server_edge(std::size_t i)
|
||||
{
|
||||
if (i >= server.size())
|
||||
throw std::out_of_range("memory_star server");
|
||||
return edge_mesh{{&server[i]}};
|
||||
}
|
||||
};
|
||||
|
||||
/// @brief Build an in-process client↔N-server star.
|
||||
/// @param slot_bytes Round widths shared by every edge (same schedule shape).
|
||||
inline memory_star make_memory_star(std::size_t n_servers, std::size_t count,
|
||||
std::vector<std::size_t> slot_bytes)
|
||||
{
|
||||
if (n_servers < 2)
|
||||
throw std::invalid_argument("make_memory_star needs >= 2 servers");
|
||||
memory_star star;
|
||||
star.servers = n_servers;
|
||||
star.hubs.reserve(n_servers);
|
||||
star.client.reserve(n_servers);
|
||||
star.server.reserve(n_servers);
|
||||
for (std::size_t i = 0; i < n_servers; ++i)
|
||||
{
|
||||
auto hub = std::make_shared<memory_sink_hub>(count, slot_bytes);
|
||||
star.hubs.push_back(hub);
|
||||
star.client.emplace_back(hub, true);
|
||||
star.server.emplace_back(hub, false);
|
||||
}
|
||||
return star;
|
||||
}
|
||||
|
||||
/// @brief Fully connected memory clique of `n` roles (every unordered pair).
|
||||
/// @details Edge id for ordered pair (a,b) with a<b is the combinatorial
|
||||
/// index; both directions share one hub (a is side_a).
|
||||
struct memory_clique
|
||||
{
|
||||
std::size_t roles = 0;
|
||||
std::vector<std::shared_ptr<memory_sink_hub>> hubs;
|
||||
/// hubs[edge], ends[edge].first = lower role, .second = higher role
|
||||
std::vector<std::pair<memory_sink, memory_sink>> ends;
|
||||
|
||||
static edge_id pair_edge(std::size_t a, std::size_t b, std::size_t n)
|
||||
{
|
||||
if (a == b || a >= n || b >= n)
|
||||
throw std::invalid_argument("memory_clique pair");
|
||||
if (a > b)
|
||||
std::swap(a, b);
|
||||
// Index among pairs (i,j) with i<j.
|
||||
edge_id e = 0;
|
||||
for (std::size_t i = 0; i < a; ++i)
|
||||
e = static_cast<edge_id>(e + (n - 1 - i));
|
||||
e = static_cast<edge_id>(e + (b - a - 1));
|
||||
return e;
|
||||
}
|
||||
|
||||
RoundSink & end(std::size_t role, std::size_t peer)
|
||||
{
|
||||
const auto e = pair_edge(role, peer, roles);
|
||||
if (role < peer)
|
||||
return ends[e].first;
|
||||
return ends[e].second;
|
||||
}
|
||||
|
||||
edge_mesh mesh_for(std::size_t role)
|
||||
{
|
||||
edge_mesh m;
|
||||
m.sinks.assign(ends.size(), nullptr);
|
||||
for (std::size_t p = 0; p < roles; ++p)
|
||||
{
|
||||
if (p == role)
|
||||
continue;
|
||||
m.sinks[pair_edge(role, p, roles)] = &end(role, p);
|
||||
}
|
||||
return m;
|
||||
}
|
||||
};
|
||||
|
||||
inline memory_clique make_memory_clique(std::size_t n_roles, std::size_t count,
|
||||
std::vector<std::size_t> slot_bytes)
|
||||
{
|
||||
if (n_roles < 2)
|
||||
throw std::invalid_argument("make_memory_clique needs >= 2 roles");
|
||||
memory_clique c;
|
||||
c.roles = n_roles;
|
||||
const std::size_t n_edges = n_roles * (n_roles - 1) / 2;
|
||||
c.hubs.reserve(n_edges);
|
||||
c.ends.reserve(n_edges);
|
||||
for (std::size_t e = 0; e < n_edges; ++e)
|
||||
{
|
||||
auto hub = std::make_shared<memory_sink_hub>(count, slot_bytes);
|
||||
c.hubs.push_back(hub);
|
||||
c.ends.emplace_back(memory_sink(hub, true), memory_sink(hub, false));
|
||||
}
|
||||
return c;
|
||||
}
|
||||
|
||||
} // namespace net
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_NET_EDGE_MESH_HPP__
|
||||
378
include/dpf/net/identity.hpp
Normal file
378
include/dpf/net/identity.hpp
Normal file
|
|
@ -0,0 +1,378 @@
|
|||
/// @file dpf/net/identity.hpp
|
||||
/// @brief Party identity keys (Ed25519) and their text and file forms.
|
||||
/// @details A party is identified by a raw 32-byte Ed25519 public key, written
|
||||
/// as 44 characters of base64 (the same shape as a WireGuard key). A
|
||||
/// key file holds the 32-byte private seed as one base64 line and is
|
||||
/// created mode 0600. TLS needs a certificate, so `identity` also
|
||||
/// carries a self-signed certificate generated in memory from the
|
||||
/// key; peers check the key inside it, never the certificate fields.
|
||||
/// `development()` is a fixed, publicly known identity: it keeps the
|
||||
/// client path working with no configuration and provides no security.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_NET_IDENTITY_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_NET_IDENTITY_HPP__
|
||||
|
||||
#include <algorithm>
|
||||
#include <array>
|
||||
#include <cerrno>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <fstream>
|
||||
#include <memory>
|
||||
#include <stdexcept>
|
||||
#include <string>
|
||||
#include <utility>
|
||||
|
||||
#include <fcntl.h>
|
||||
#include <sys/stat.h>
|
||||
#include <unistd.h>
|
||||
|
||||
#ifndef DPF_HAS_OPENSSL
|
||||
#if defined(__has_include)
|
||||
#if __has_include(<openssl/ssl.h>)
|
||||
#define DPF_HAS_OPENSSL 1
|
||||
#endif
|
||||
#endif
|
||||
#endif
|
||||
#ifndef DPF_HAS_OPENSSL
|
||||
#define DPF_HAS_OPENSSL 0
|
||||
#endif
|
||||
|
||||
#if DPF_HAS_OPENSSL
|
||||
#include <openssl/err.h>
|
||||
#include <openssl/evp.h>
|
||||
#include <openssl/rand.h>
|
||||
#include <openssl/x509.h>
|
||||
#endif
|
||||
|
||||
#include "dpf/log.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace net
|
||||
{
|
||||
namespace detail
|
||||
{
|
||||
|
||||
inline std::string base64_encode(const std::uint8_t * p, std::size_t n)
|
||||
{
|
||||
static const char tab[] =
|
||||
"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
|
||||
std::string out;
|
||||
out.reserve((n + 2) / 3 * 4);
|
||||
for (std::size_t i = 0; i < n; i += 3)
|
||||
{
|
||||
const std::uint32_t a = p[i];
|
||||
const std::uint32_t b = i + 1 < n ? p[i + 1] : 0;
|
||||
const std::uint32_t c = i + 2 < n ? p[i + 2] : 0;
|
||||
const std::uint32_t v = (a << 16) | (b << 8) | c;
|
||||
out += tab[(v >> 18) & 63];
|
||||
out += tab[(v >> 12) & 63];
|
||||
out += i + 1 < n ? tab[(v >> 6) & 63] : '=';
|
||||
out += i + 2 < n ? tab[v & 63] : '=';
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
/// @brief Strict base64 (standard alphabet, padded). Returns false on any
|
||||
/// malformed input.
|
||||
inline bool base64_decode(const std::string & s, std::string & out)
|
||||
{
|
||||
auto val = [](char ch) -> int {
|
||||
if (ch >= 'A' && ch <= 'Z')
|
||||
return ch - 'A';
|
||||
if (ch >= 'a' && ch <= 'z')
|
||||
return ch - 'a' + 26;
|
||||
if (ch >= '0' && ch <= '9')
|
||||
return ch - '0' + 52;
|
||||
if (ch == '+')
|
||||
return 62;
|
||||
if (ch == '/')
|
||||
return 63;
|
||||
return -1;
|
||||
};
|
||||
out.clear();
|
||||
if (s.size() % 4 != 0)
|
||||
return false;
|
||||
for (std::size_t i = 0; i < s.size(); i += 4)
|
||||
{
|
||||
const bool last = i + 4 == s.size();
|
||||
int v[4];
|
||||
for (int k = 0; k < 4; ++k)
|
||||
{
|
||||
const char ch = s[i + k];
|
||||
if (ch == '=' && last && k >= 2)
|
||||
v[k] = -2;
|
||||
else
|
||||
v[k] = val(ch);
|
||||
if (v[k] == -1)
|
||||
return false;
|
||||
}
|
||||
if (v[0] < 0 || v[1] < 0 || (v[2] == -2 && v[3] != -2))
|
||||
return false;
|
||||
const std::uint32_t x = (static_cast<std::uint32_t>(v[0]) << 18)
|
||||
| (static_cast<std::uint32_t>(v[1]) << 12)
|
||||
| (static_cast<std::uint32_t>(v[2] < 0 ? 0 : v[2]) << 6)
|
||||
| static_cast<std::uint32_t>(v[3] < 0 ? 0 : v[3]);
|
||||
out += static_cast<char>((x >> 16) & 0xff);
|
||||
if (v[2] >= 0)
|
||||
out += static_cast<char>((x >> 8) & 0xff);
|
||||
if (v[3] >= 0)
|
||||
out += static_cast<char>(x & 0xff);
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
/// @brief First line of `path` that is neither blank nor a `#` comment.
|
||||
inline std::string first_key_line(const std::string & path, const char * what)
|
||||
{
|
||||
std::ifstream in(path);
|
||||
if (!in)
|
||||
throw std::runtime_error(std::string(what) + ": cannot read '" + path + "'");
|
||||
std::string line;
|
||||
while (std::getline(in, line))
|
||||
{
|
||||
while (!line.empty()
|
||||
&& (line.back() == '\r' || line.back() == ' ' || line.back() == '\t'))
|
||||
line.pop_back();
|
||||
std::size_t b = 0;
|
||||
while (b < line.size() && (line[b] == ' ' || line[b] == '\t'))
|
||||
++b;
|
||||
line.erase(0, b);
|
||||
if (!line.empty() && line[0] != '#')
|
||||
return line;
|
||||
}
|
||||
throw std::runtime_error(std::string(what) + ": '" + path + "' holds no key");
|
||||
}
|
||||
|
||||
#if DPF_HAS_OPENSSL
|
||||
inline std::string openssl_error(const char * what)
|
||||
{
|
||||
std::string out = what;
|
||||
unsigned long e = 0;
|
||||
bool first = true;
|
||||
while ((e = ERR_get_error()) != 0)
|
||||
{
|
||||
char buf[256];
|
||||
ERR_error_string_n(e, buf, sizeof(buf));
|
||||
out += first ? ": " : "; ";
|
||||
out += buf;
|
||||
first = false;
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
struct pkey_free
|
||||
{
|
||||
void operator()(EVP_PKEY * k) const noexcept { EVP_PKEY_free(k); }
|
||||
};
|
||||
struct x509_free
|
||||
{
|
||||
void operator()(X509 * x) const noexcept { X509_free(x); }
|
||||
};
|
||||
#endif
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/// @brief A party's raw Ed25519 public key.
|
||||
struct public_key
|
||||
{
|
||||
static constexpr std::size_t size = 32;
|
||||
std::array<std::uint8_t, size> bytes{};
|
||||
|
||||
/// @brief 44 characters of base64.
|
||||
std::string base64() const { return detail::base64_encode(bytes.data(), size); }
|
||||
|
||||
/// @brief Parse base64, or read a public-key file given as `file:PATH`.
|
||||
static public_key parse(const std::string & text)
|
||||
{
|
||||
std::string s = text;
|
||||
if (s.rfind("file:", 0) == 0)
|
||||
s = detail::first_key_line(s.substr(5), "public key");
|
||||
std::string raw;
|
||||
if (!detail::base64_decode(s, raw) || raw.size() != size)
|
||||
throw std::invalid_argument("public key must be 32 bytes of base64 "
|
||||
"(44 characters), got '" + s + "'");
|
||||
public_key k;
|
||||
std::memcpy(k.bytes.data(), raw.data(), size);
|
||||
return k;
|
||||
}
|
||||
|
||||
friend bool operator==(const public_key & a, const public_key & b) noexcept
|
||||
{
|
||||
return a.bytes == b.bytes;
|
||||
}
|
||||
friend bool operator!=(const public_key & a, const public_key & b) noexcept
|
||||
{
|
||||
return !(a == b);
|
||||
}
|
||||
};
|
||||
|
||||
/// @brief A party's private key plus the self-signed certificate TLS presents.
|
||||
/// @details Copies share the key. Without OpenSSL every constructor throws.
|
||||
class identity
|
||||
{
|
||||
public:
|
||||
/// @brief A fresh random key.
|
||||
static identity generate()
|
||||
{
|
||||
#if DPF_HAS_OPENSSL
|
||||
std::uint8_t seed[32];
|
||||
if (RAND_bytes(seed, sizeof(seed)) != 1)
|
||||
throw std::runtime_error(detail::openssl_error("identity: RAND_bytes"));
|
||||
auto id = from_seed(seed);
|
||||
OPENSSL_cleanse(seed, sizeof(seed));
|
||||
return id;
|
||||
#else
|
||||
throw std::logic_error("identity: built without OpenSSL");
|
||||
#endif
|
||||
}
|
||||
|
||||
/// @brief The key whose private seed is `seed`.
|
||||
static identity from_seed(const std::uint8_t * seed, const char * cn = "libdpf party")
|
||||
{
|
||||
#if DPF_HAS_OPENSSL
|
||||
identity id;
|
||||
EVP_PKEY * k = EVP_PKEY_new_raw_private_key(EVP_PKEY_ED25519, nullptr, seed, 32);
|
||||
if (k == nullptr)
|
||||
throw std::runtime_error(detail::openssl_error("identity: bad Ed25519 seed"));
|
||||
id.key_.reset(k, detail::pkey_free{});
|
||||
std::size_t n = public_key::size;
|
||||
if (EVP_PKEY_get_raw_public_key(k, id.pub_.bytes.data(), &n) != 1
|
||||
|| n != public_key::size)
|
||||
throw std::runtime_error(detail::openssl_error("identity: public key"));
|
||||
id.cert_ = self_signed(k, cn);
|
||||
return id;
|
||||
#else
|
||||
(void)seed;
|
||||
(void)cn;
|
||||
throw std::logic_error("identity: built without OpenSSL");
|
||||
#endif
|
||||
}
|
||||
|
||||
/// @brief Read a key file (one base64 line holding the 32-byte seed).
|
||||
/// @details Warns when the file is readable by group or others.
|
||||
static identity load(const std::string & path)
|
||||
{
|
||||
std::string raw;
|
||||
const auto line = detail::first_key_line(path, "identity");
|
||||
if (!detail::base64_decode(line, raw) || raw.size() != 32)
|
||||
throw std::invalid_argument("identity: '" + path
|
||||
+ "' is not a key file (expected 32 bytes of base64)");
|
||||
struct stat st{};
|
||||
if (::stat(path.c_str(), &st) == 0 && (st.st_mode & 077) != 0)
|
||||
DPF_LOG(warning, "security.key_file_mode").kv("path", path)
|
||||
.kv("detail", "identity key file is readable by group or others; "
|
||||
"chmod 600 it");
|
||||
auto id = from_seed(reinterpret_cast<const std::uint8_t *>(raw.data()));
|
||||
std::fill(raw.begin(), raw.end(), '\0');
|
||||
return id;
|
||||
}
|
||||
|
||||
/// @brief The fixed development identity. Its private key is in this
|
||||
/// source file, so it authenticates nothing.
|
||||
static const identity & development()
|
||||
{
|
||||
static const identity dev = [] {
|
||||
#if DPF_HAS_OPENSSL
|
||||
static const char phrase[] =
|
||||
"libdpf development identity (public; provides no security)";
|
||||
std::uint8_t seed[32];
|
||||
unsigned int n = sizeof(seed);
|
||||
if (EVP_Digest(phrase, sizeof(phrase) - 1, seed, &n, EVP_sha256(), nullptr)
|
||||
!= 1)
|
||||
throw std::runtime_error(detail::openssl_error("identity: digest"));
|
||||
auto id = from_seed(seed, "libdpf development certificate (no security)");
|
||||
id.development_ = true;
|
||||
return id;
|
||||
#else
|
||||
return identity();
|
||||
#endif
|
||||
}();
|
||||
if (!dev.key_)
|
||||
throw std::logic_error("identity: built without OpenSSL");
|
||||
return dev;
|
||||
}
|
||||
|
||||
/// @brief Write the private seed to a new file with mode 0600. Refuses to
|
||||
/// replace an existing file unless `overwrite`.
|
||||
void save(const std::string & path, bool overwrite = false) const
|
||||
{
|
||||
#if DPF_HAS_OPENSSL
|
||||
std::uint8_t seed[32];
|
||||
std::size_t n = sizeof(seed);
|
||||
if (!key_ || EVP_PKEY_get_raw_private_key(key_.get(), seed, &n) != 1 || n != 32)
|
||||
throw std::runtime_error(detail::openssl_error("identity: private key"));
|
||||
const std::string text = "# libdpf identity key (private; keep mode 600)\n"
|
||||
+ detail::base64_encode(seed, sizeof(seed)) + "\n";
|
||||
OPENSSL_cleanse(seed, sizeof(seed));
|
||||
const int flags = O_WRONLY | O_CREAT | O_CLOEXEC | (overwrite ? O_TRUNC : O_EXCL);
|
||||
const int fd = ::open(path.c_str(), flags, 0600);
|
||||
if (fd < 0)
|
||||
throw std::runtime_error("identity: cannot create '" + path + "': "
|
||||
+ std::strerror(errno));
|
||||
const bool ok = ::fchmod(fd, 0600) == 0
|
||||
&& ::write(fd, text.data(), text.size()) == static_cast<ssize_t>(text.size());
|
||||
::close(fd);
|
||||
if (!ok)
|
||||
throw std::runtime_error("identity: cannot write '" + path + "'");
|
||||
#else
|
||||
(void)path;
|
||||
(void)overwrite;
|
||||
throw std::logic_error("identity: built without OpenSSL");
|
||||
#endif
|
||||
}
|
||||
|
||||
const public_key & key() const noexcept { return pub_; }
|
||||
bool is_development() const noexcept { return development_; }
|
||||
|
||||
#if DPF_HAS_OPENSSL
|
||||
EVP_PKEY * pkey() const noexcept { return key_.get(); }
|
||||
X509 * cert() const noexcept { return cert_.get(); }
|
||||
#endif
|
||||
|
||||
private:
|
||||
identity() = default;
|
||||
|
||||
#if DPF_HAS_OPENSSL
|
||||
static std::shared_ptr<X509> self_signed(EVP_PKEY * k, const char * cn)
|
||||
{
|
||||
std::shared_ptr<X509> x(X509_new(), detail::x509_free{});
|
||||
if (!x)
|
||||
throw std::runtime_error(detail::openssl_error("identity: X509_new"));
|
||||
std::uint8_t serial[8];
|
||||
if (RAND_bytes(serial, sizeof(serial)) != 1)
|
||||
throw std::runtime_error(detail::openssl_error("identity: serial"));
|
||||
std::uint64_t s = 0;
|
||||
for (auto b : serial)
|
||||
s = (s << 8) | b;
|
||||
s &= 0x7fffffffffffffffull;
|
||||
bool ok = X509_set_version(x.get(), 2) == 1
|
||||
&& ASN1_INTEGER_set_uint64(X509_get_serialNumber(x.get()), s) == 1
|
||||
&& X509_gmtime_adj(X509_getm_notBefore(x.get()), -86400) != nullptr
|
||||
&& X509_time_adj_ex(X509_getm_notAfter(x.get()), 36500, 0, nullptr) != nullptr
|
||||
&& X509_set_pubkey(x.get(), k) == 1;
|
||||
X509_NAME * name = X509_get_subject_name(x.get());
|
||||
ok = ok && name != nullptr
|
||||
&& X509_NAME_add_entry_by_txt(name, "CN", MBSTRING_ASC,
|
||||
reinterpret_cast<const unsigned char *>(cn), -1, -1, 0)
|
||||
== 1
|
||||
&& X509_set_issuer_name(x.get(), name) == 1
|
||||
&& X509_sign(x.get(), k, nullptr) > 0;
|
||||
if (!ok)
|
||||
throw std::runtime_error(detail::openssl_error("identity: certificate"));
|
||||
return x;
|
||||
}
|
||||
|
||||
std::shared_ptr<EVP_PKEY> key_;
|
||||
std::shared_ptr<X509> cert_;
|
||||
#else
|
||||
std::shared_ptr<void> key_;
|
||||
#endif
|
||||
public_key pub_{};
|
||||
bool development_ = false;
|
||||
};
|
||||
|
||||
} // namespace net
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_NET_IDENTITY_HPP__
|
||||
88
include/dpf/net/io_pool.hpp
Normal file
88
include/dpf/net/io_pool.hpp
Normal file
|
|
@ -0,0 +1,88 @@
|
|||
/// @file dpf/net/io_pool.hpp
|
||||
/// @brief Socket completion threads and a separate compute pool.
|
||||
/// @details `context()` is run by `io_threads` workers (0 = hardware
|
||||
/// concurrency). `post_compute` runs on `compute_threads` workers
|
||||
/// (0 = same as io), so a DPF evaluation never occupies a thread that
|
||||
/// should be completing reads and writes. Shutdown stops socket I/O
|
||||
/// first, then joins compute.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_NET_IO_POOL_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_NET_IO_POOL_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <thread>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/net/asio_ns.hpp"
|
||||
#include <asio/thread_pool.hpp>
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace net
|
||||
{
|
||||
|
||||
class io_pool
|
||||
{
|
||||
public:
|
||||
explicit io_pool(std::size_t io_threads = 0, std::size_t compute_threads = 0)
|
||||
: work_(asio::make_work_guard(io_)),
|
||||
compute_n_(pick(compute_threads == 0 ? io_threads : compute_threads)),
|
||||
compute_(compute_n_)
|
||||
{
|
||||
const std::size_t n = pick(io_threads);
|
||||
threads_.reserve(n);
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
threads_.emplace_back([this] { io_.run(); });
|
||||
}
|
||||
|
||||
io_pool(const io_pool &) = delete;
|
||||
io_pool & operator=(const io_pool &) = delete;
|
||||
|
||||
~io_pool()
|
||||
{
|
||||
work_.reset();
|
||||
io_.stop();
|
||||
for (auto & t : threads_)
|
||||
if (t.joinable())
|
||||
t.join();
|
||||
compute_.join();
|
||||
}
|
||||
|
||||
asio::io_context & context() noexcept { return io_; }
|
||||
std::size_t size() const noexcept { return threads_.size(); }
|
||||
std::size_t compute_size() const noexcept { return compute_n_; }
|
||||
|
||||
/// @brief Run `fn` on a socket worker.
|
||||
template <typename Fn>
|
||||
void post(Fn && fn)
|
||||
{
|
||||
asio::post(io_, std::forward<Fn>(fn));
|
||||
}
|
||||
|
||||
/// @brief Run `fn` on the compute pool, leaving socket workers free.
|
||||
template <typename Fn>
|
||||
void post_compute(Fn && fn)
|
||||
{
|
||||
asio::post(compute_, std::forward<Fn>(fn));
|
||||
}
|
||||
|
||||
private:
|
||||
static std::size_t pick(std::size_t n)
|
||||
{
|
||||
if (n != 0)
|
||||
return n;
|
||||
const auto hw = std::thread::hardware_concurrency();
|
||||
return hw == 0 ? 1 : hw;
|
||||
}
|
||||
|
||||
asio::io_context io_;
|
||||
asio::executor_work_guard<asio::io_context::executor_type> work_;
|
||||
std::size_t compute_n_ = 1;
|
||||
asio::thread_pool compute_;
|
||||
std::vector<std::thread> threads_;
|
||||
};
|
||||
|
||||
} // namespace net
|
||||
} // namespace dpf
|
||||
|
||||
#endif
|
||||
161
include/dpf/net/link_log.hpp
Normal file
161
include/dpf/net/link_log.hpp
Normal file
|
|
@ -0,0 +1,161 @@
|
|||
/// @file dpf/net/link_log.hpp
|
||||
/// @brief Run-log records for listeners and established party links.
|
||||
/// @details `log_link_up` is called once per socket after the handshake and
|
||||
/// after the stream array has adopted and tuned it, so the socket
|
||||
/// options it reads back are the ones the kernel applied (Linux
|
||||
/// doubles `SO_SNDBUF`/`SO_RCVBUF` and clamps them to `wmem_max` and
|
||||
/// `rmem_max`). `TCP_INFO` at that point carries the kernel's RTT
|
||||
/// estimate from the connection setup and the handshake exchange.
|
||||
/// Encrypted links record the TLS version and cipher, how this side
|
||||
/// authenticated the peer (`auth=key` or `none`), the peer's key, and
|
||||
/// whether the peer authenticated this side. With encryption off the
|
||||
/// record says `auth=none encryption=none`, and the first plaintext
|
||||
/// link to an address off this host also raises one warning.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_NET_LINK_LOG_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_NET_LINK_LOG_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <string>
|
||||
|
||||
#include <arpa/inet.h>
|
||||
#include <netinet/in.h>
|
||||
#include <netinet/tcp.h>
|
||||
#include <sys/socket.h>
|
||||
#include <sys/un.h>
|
||||
|
||||
#include "dpf/log.hpp"
|
||||
#include "dpf/net/policy.hpp"
|
||||
#include "dpf/net/security.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace net
|
||||
{
|
||||
namespace detail
|
||||
{
|
||||
|
||||
inline std::string sockaddr_text(const sockaddr_storage & ss)
|
||||
{
|
||||
char host[INET6_ADDRSTRLEN] = {};
|
||||
if (ss.ss_family == AF_INET)
|
||||
{
|
||||
const auto & a = reinterpret_cast<const sockaddr_in &>(ss);
|
||||
if (::inet_ntop(AF_INET, &a.sin_addr, host, sizeof(host)) == nullptr)
|
||||
return "unknown";
|
||||
return std::string(host) + ":" + std::to_string(ntohs(a.sin_port));
|
||||
}
|
||||
if (ss.ss_family == AF_INET6)
|
||||
{
|
||||
const auto & a = reinterpret_cast<const sockaddr_in6 &>(ss);
|
||||
if (::inet_ntop(AF_INET6, &a.sin6_addr, host, sizeof(host)) == nullptr)
|
||||
return "unknown";
|
||||
return "[" + std::string(host) + "]:" + std::to_string(ntohs(a.sin6_port));
|
||||
}
|
||||
if (ss.ss_family == AF_UNIX)
|
||||
{
|
||||
const auto & a = reinterpret_cast<const sockaddr_un &>(ss);
|
||||
return std::string("unix:") + (a.sun_path[0] != '\0' ? a.sun_path : "(unnamed)");
|
||||
}
|
||||
return "unknown";
|
||||
}
|
||||
|
||||
inline bool loopback(const sockaddr_storage & ss)
|
||||
{
|
||||
if (ss.ss_family == AF_INET)
|
||||
return (ntohl(reinterpret_cast<const sockaddr_in &>(ss).sin_addr.s_addr) >> 24)
|
||||
== 127u;
|
||||
if (ss.ss_family == AF_INET6)
|
||||
{
|
||||
const auto & a = reinterpret_cast<const sockaddr_in6 &>(ss).sin6_addr;
|
||||
if (IN6_IS_ADDR_LOOPBACK(&a))
|
||||
return true;
|
||||
return IN6_IS_ADDR_V4MAPPED(&a) && a.s6_addr[12] == 127;
|
||||
}
|
||||
return ss.ss_family == AF_UNIX;
|
||||
}
|
||||
|
||||
inline int int_opt(int fd, int level, int name)
|
||||
{
|
||||
int v = -1;
|
||||
socklen_t len = sizeof(v);
|
||||
if (::getsockopt(fd, level, name, &v, &len) != 0)
|
||||
return -1;
|
||||
return v;
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/// @brief Record a listener: this process accepts any address on `port`.
|
||||
/// @param encrypted whether connections on it must complete TLS 1.3 first
|
||||
inline void log_listen(unsigned short port, bool sctp, bool encrypted = false)
|
||||
{
|
||||
DPF_LOG(info, "listen").kv("addr", "0.0.0.0").kv("port", port)
|
||||
.kv("tcp", true).kv("sctp", sctp)
|
||||
.kv("encryption", encrypted ? "tls1.3" : "none");
|
||||
}
|
||||
|
||||
/// @brief Record one established socket of a party link.
|
||||
/// @param how `accept` or `connect`
|
||||
/// @param peer the other end's role (`p1`, `dealer`, ...)
|
||||
/// @param lane which socket of a `parallel` link (0 otherwise)
|
||||
/// @param sec how the link was secured (null or unencrypted: plaintext)
|
||||
inline void log_link_up(const char * how, const std::string & peer, transport kind,
|
||||
std::size_t lanes, std::uint32_t lane, std::uint32_t epoch, int fd,
|
||||
const socket_options & requested, const link_security * sec = nullptr)
|
||||
{
|
||||
const bool encrypted = sec != nullptr && sec->encrypted;
|
||||
if (!log::enabled(log::level::info) || fd < 0)
|
||||
return;
|
||||
sockaddr_storage local{};
|
||||
sockaddr_storage remote{};
|
||||
socklen_t local_len = sizeof(local);
|
||||
socklen_t remote_len = sizeof(remote);
|
||||
const bool have_local =
|
||||
::getsockname(fd, reinterpret_cast<sockaddr *>(&local), &local_len) == 0;
|
||||
const bool have_remote =
|
||||
::getpeername(fd, reinterpret_cast<sockaddr *>(&remote), &remote_len) == 0;
|
||||
{
|
||||
log::record rec(log::level::info, "link.up");
|
||||
rec.kv("how", how).kv("peer", peer).kv("transport", transport_name(kind))
|
||||
.kv("lanes", lanes).kv("lane", lane).kv("epoch", epoch)
|
||||
.kv("local", have_local ? detail::sockaddr_text(local) : std::string("unknown"))
|
||||
.kv("remote", have_remote ? detail::sockaddr_text(remote) : std::string("unknown"))
|
||||
.kv("auth", encrypted ? sec->peer_auth : std::string("none"))
|
||||
.kv("encryption",
|
||||
encrypted ? sec->protocol + "/" + sec->cipher : std::string("none"));
|
||||
if (encrypted)
|
||||
rec.kv("peer_key", sec->peer_key ? sec->peer_key->base64() : std::string("none"))
|
||||
.kv("peer_verified_us", sec->peer_verified_us);
|
||||
if (kind != transport::sctp)
|
||||
{
|
||||
rec.kv("nodelay", detail::int_opt(fd, IPPROTO_TCP, TCP_NODELAY))
|
||||
.kv("quickack_req", requested.quickack)
|
||||
.kv("keepalive", detail::int_opt(fd, SOL_SOCKET, SO_KEEPALIVE))
|
||||
.kv("sndbuf_req", requested.send_buffer)
|
||||
.kv("sndbuf", detail::int_opt(fd, SOL_SOCKET, SO_SNDBUF))
|
||||
.kv("rcvbuf_req", requested.recv_buffer)
|
||||
.kv("rcvbuf", detail::int_opt(fd, SOL_SOCKET, SO_RCVBUF));
|
||||
#if defined(TCP_INFO)
|
||||
tcp_info ti{};
|
||||
socklen_t ti_len = sizeof(ti);
|
||||
if (::getsockopt(fd, IPPROTO_TCP, TCP_INFO, &ti, &ti_len) == 0)
|
||||
rec.kv("rtt_us", ti.tcpi_rtt).kv("rttvar_us", ti.tcpi_rttvar)
|
||||
.kv("pmtu", ti.tcpi_pmtu).kv("snd_mss", ti.tcpi_snd_mss)
|
||||
.kv("snd_cwnd", ti.tcpi_snd_cwnd)
|
||||
.kv("retrans", ti.tcpi_total_retrans);
|
||||
#endif
|
||||
}
|
||||
}
|
||||
if (!encrypted && have_remote && !detail::loopback(remote)
|
||||
&& log::first_time("net.plaintext_remote"))
|
||||
DPF_LOG(warning, "link.plaintext").kv("remote", detail::sockaddr_text(remote))
|
||||
.kv("detail", "encryption is off: this party link is unauthenticated "
|
||||
"and unencrypted, and the handshake's party id is not verified");
|
||||
}
|
||||
|
||||
} // namespace net
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_NET_LINK_LOG_HPP__
|
||||
179
include/dpf/net/memory_sink.hpp
Normal file
179
include/dpf/net/memory_sink.hpp
Normal file
|
|
@ -0,0 +1,179 @@
|
|||
/// @file dpf/net/memory_sink.hpp
|
||||
/// @brief In-process paired RoundSink for correctness tests.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_NET_MEMORY_SINK_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_NET_MEMORY_SINK_HPP__
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <memory>
|
||||
#include <mutex>
|
||||
#include <stdexcept>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/net/round_sink.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace net
|
||||
{
|
||||
|
||||
/// @brief Shared state between the two ends of a memory sink pair.
|
||||
struct memory_sink_hub
|
||||
{
|
||||
std::size_t count = 0;
|
||||
std::vector<std::size_t> slot_bytes;
|
||||
std::vector<round_window> a;
|
||||
std::vector<round_window> b;
|
||||
mutable std::mutex mu;
|
||||
|
||||
explicit memory_sink_hub(std::size_t n, std::vector<std::size_t> slots)
|
||||
: count(n), slot_bytes(std::move(slots))
|
||||
{
|
||||
a.reserve(slot_bytes.size());
|
||||
b.reserve(slot_bytes.size());
|
||||
for (std::size_t sb : slot_bytes)
|
||||
{
|
||||
a.emplace_back(count, sb);
|
||||
b.emplace_back(count, sb);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
/// @brief One end of a memory-paired RoundSink.
|
||||
class memory_sink : public RoundSink
|
||||
{
|
||||
public:
|
||||
memory_sink(std::shared_ptr<memory_sink_hub> hub, bool side_a)
|
||||
: hub_(std::move(hub)), side_a_(side_a)
|
||||
{
|
||||
if (!hub_)
|
||||
throw std::invalid_argument("memory_sink needs a hub");
|
||||
}
|
||||
|
||||
std::size_t count() const noexcept override { return hub_->count; }
|
||||
std::size_t rounds() const noexcept override
|
||||
{
|
||||
return hub_->slot_bytes.size();
|
||||
}
|
||||
|
||||
std::size_t slot_bytes(std::uint16_t round) const override
|
||||
{
|
||||
if (round >= hub_->slot_bytes.size())
|
||||
throw std::out_of_range("memory_sink round");
|
||||
return hub_->slot_bytes[round];
|
||||
}
|
||||
|
||||
void submit(std::uint16_t round, std::size_t index,
|
||||
const std::uint8_t * bytes, std::size_t n) override
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(hub_->mu);
|
||||
mine(round).submit(index, bytes, n);
|
||||
}
|
||||
|
||||
bool peer_ready(std::uint16_t round, std::size_t index) const override
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(hub_->mu);
|
||||
return mine(round).peer_ready(index);
|
||||
}
|
||||
|
||||
void read_peer(std::uint16_t round, std::size_t index, std::uint8_t * out,
|
||||
std::size_t n) const override
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(hub_->mu);
|
||||
mine(round).read_peer(index, out, n);
|
||||
}
|
||||
|
||||
void flush() override
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(hub_->mu);
|
||||
flush_unlocked();
|
||||
}
|
||||
|
||||
void flush_round(std::uint16_t round) override
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(hub_->mu);
|
||||
flush_round_unlocked(round);
|
||||
}
|
||||
|
||||
void poll() override {}
|
||||
|
||||
private:
|
||||
void flush_unlocked()
|
||||
{
|
||||
for (std::uint16_t r = 0; r < hub_->slot_bytes.size(); ++r)
|
||||
{
|
||||
std::size_t begin = 0;
|
||||
std::size_t n = 0;
|
||||
mine(r).pending_out(begin, n);
|
||||
if (n != 0)
|
||||
{
|
||||
flush_round_unlocked(r);
|
||||
continue;
|
||||
}
|
||||
peer(r).pending_out(begin, n);
|
||||
if (n != 0)
|
||||
flush_round_unlocked(r);
|
||||
}
|
||||
}
|
||||
|
||||
void flush_round_unlocked(std::uint16_t round)
|
||||
{
|
||||
if (round >= hub_->slot_bytes.size())
|
||||
throw std::out_of_range("memory_sink flush_round");
|
||||
auto & local = mine(round);
|
||||
auto & remote = peer(round);
|
||||
std::size_t begin_l = 0;
|
||||
std::size_t n_l = 0;
|
||||
const std::uint8_t * pend_l = local.pending_out(begin_l, n_l);
|
||||
std::size_t begin_r = 0;
|
||||
std::size_t n_r = 0;
|
||||
const std::uint8_t * pend_r = remote.pending_out(begin_r, n_r);
|
||||
if (n_l != 0)
|
||||
{
|
||||
remote.accept_peer_at(begin_l, pend_l, n_l);
|
||||
local.mark_flushed(n_l);
|
||||
}
|
||||
if (n_r != 0)
|
||||
{
|
||||
local.accept_peer_at(begin_r, pend_r, n_r);
|
||||
remote.mark_flushed(n_r);
|
||||
}
|
||||
}
|
||||
round_window & mine(std::uint16_t round)
|
||||
{
|
||||
if (round >= hub_->slot_bytes.size())
|
||||
throw std::out_of_range("memory_sink round");
|
||||
return side_a_ ? hub_->a[round] : hub_->b[round];
|
||||
}
|
||||
|
||||
const round_window & mine(std::uint16_t round) const
|
||||
{
|
||||
if (round >= hub_->slot_bytes.size())
|
||||
throw std::out_of_range("memory_sink round");
|
||||
return side_a_ ? hub_->a[round] : hub_->b[round];
|
||||
}
|
||||
|
||||
round_window & peer(std::uint16_t round)
|
||||
{
|
||||
if (round >= hub_->slot_bytes.size())
|
||||
throw std::out_of_range("memory_sink round");
|
||||
return side_a_ ? hub_->b[round] : hub_->a[round];
|
||||
}
|
||||
|
||||
std::shared_ptr<memory_sink_hub> hub_;
|
||||
bool side_a_ = true;
|
||||
};
|
||||
|
||||
/// @brief Build a connected pair of memory sinks that share one hub.
|
||||
inline std::pair<memory_sink, memory_sink> make_memory_sink_pair(
|
||||
std::size_t count, std::vector<std::size_t> slot_bytes)
|
||||
{
|
||||
auto hub = std::make_shared<memory_sink_hub>(count, std::move(slot_bytes));
|
||||
return {memory_sink(hub, true), memory_sink(hub, false)};
|
||||
}
|
||||
|
||||
} // namespace net
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_NET_MEMORY_SINK_HPP__
|
||||
28
include/dpf/net/mesh_rendezvous.hpp
Normal file
28
include/dpf/net/mesh_rendezvous.hpp
Normal file
|
|
@ -0,0 +1,28 @@
|
|||
/// @file dpf/net/mesh_rendezvous.hpp
|
||||
/// @brief Shared in-process port table for party clique joins.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_NET_MESH_RENDEZVOUS_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_NET_MESH_RENDEZVOUS_HPP__
|
||||
|
||||
#include <atomic>
|
||||
#include <vector>
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace net
|
||||
{
|
||||
|
||||
/// @brief Shared port table: `ports[i]` is party `i`'s listen port (0 until bound).
|
||||
using mesh_ports = std::vector<std::atomic<unsigned short>>;
|
||||
|
||||
inline mesh_ports make_mesh_ports(unsigned n)
|
||||
{
|
||||
mesh_ports ports(n);
|
||||
for (auto & p : ports)
|
||||
p.store(0);
|
||||
return ports;
|
||||
}
|
||||
|
||||
} // namespace net
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_NET_MESH_RENDEZVOUS_HPP__
|
||||
269
include/dpf/net/mux_sink.hpp
Normal file
269
include/dpf/net/mux_sink.hpp
Normal file
|
|
@ -0,0 +1,269 @@
|
|||
/// @file dpf/net/mux_sink.hpp
|
||||
/// @brief RoundSink multiplexed on one framed trio channel.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_NET_MUX_SINK_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_NET_MUX_SINK_HPP__
|
||||
|
||||
#include <array>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <memory>
|
||||
#include <stdexcept>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/net/channel.hpp"
|
||||
#include "dpf/net/comm_hook.hpp"
|
||||
#include "dpf/net/round_sink.hpp"
|
||||
#include "dpf/net/trio.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace net
|
||||
{
|
||||
|
||||
#pragma pack(push, 1)
|
||||
struct round_batch_hdr
|
||||
{
|
||||
std::uint16_t round = 0;
|
||||
std::uint32_t begin = 0;
|
||||
std::uint32_t count = 0;
|
||||
};
|
||||
#pragma pack(pop)
|
||||
|
||||
/// @brief Prefix-flush RoundSink over a single duplex `channel`.
|
||||
/// @details Each flush sends one `msg::round_batch` frame per round that has
|
||||
/// a new contiguous prefix: header `(round, begin, count)` then
|
||||
/// `count * slot_bytes` payload. The peer's matching frames are
|
||||
/// read in the same flush (lower role sends first).
|
||||
class mux_sink : public RoundSink
|
||||
{
|
||||
public:
|
||||
mux_sink(channel & link, unsigned self_id, unsigned peer_id,
|
||||
std::size_t count, std::vector<std::size_t> slot_bytes)
|
||||
: link_(link),
|
||||
self_id_(self_id),
|
||||
peer_id_(peer_id),
|
||||
count_(count),
|
||||
slot_bytes_(std::move(slot_bytes))
|
||||
{
|
||||
windows_.reserve(slot_bytes_.size());
|
||||
for (std::size_t sb : slot_bytes_)
|
||||
windows_.emplace_back(count_, sb);
|
||||
}
|
||||
|
||||
std::size_t count() const noexcept override { return count_; }
|
||||
std::size_t rounds() const noexcept override { return slot_bytes_.size(); }
|
||||
|
||||
std::size_t slot_bytes(std::uint16_t round) const override
|
||||
{
|
||||
if (round >= slot_bytes_.size())
|
||||
throw std::out_of_range("mux_sink round");
|
||||
return slot_bytes_[round];
|
||||
}
|
||||
|
||||
void submit(std::uint16_t round, std::size_t index,
|
||||
const std::uint8_t * bytes, std::size_t n) override
|
||||
{
|
||||
window(round).submit(index, bytes, n);
|
||||
}
|
||||
|
||||
bool peer_ready(std::uint16_t round, std::size_t index) const override
|
||||
{
|
||||
return window(round).peer_ready(index);
|
||||
}
|
||||
|
||||
void read_peer(std::uint16_t round, std::size_t index, std::uint8_t * out,
|
||||
std::size_t n) const override
|
||||
{
|
||||
window(round).read_peer(index, out, n);
|
||||
}
|
||||
|
||||
void flush() override
|
||||
{
|
||||
std::vector<std::uint8_t> payload;
|
||||
for (std::uint16_t r = 0; r < slot_bytes_.size(); ++r)
|
||||
append_pending(r, payload);
|
||||
exchange_payload(payload);
|
||||
}
|
||||
|
||||
void flush_round(std::uint16_t round) override
|
||||
{
|
||||
if (round >= slot_bytes_.size())
|
||||
throw std::out_of_range("mux_sink flush_round");
|
||||
std::vector<std::uint8_t> payload;
|
||||
append_pending(round, payload);
|
||||
exchange_payload(payload);
|
||||
}
|
||||
|
||||
void poll() override {}
|
||||
|
||||
std::uint64_t exchanges() const noexcept { return link_exchanges_; }
|
||||
|
||||
private:
|
||||
round_window & window(std::uint16_t round)
|
||||
{
|
||||
if (round >= windows_.size())
|
||||
throw std::out_of_range("mux_sink round");
|
||||
return windows_[round];
|
||||
}
|
||||
|
||||
const round_window & window(std::uint16_t round) const
|
||||
{
|
||||
if (round >= windows_.size())
|
||||
throw std::out_of_range("mux_sink round");
|
||||
return windows_[round];
|
||||
}
|
||||
|
||||
void append_pending(std::uint16_t r, std::vector<std::uint8_t> & payload)
|
||||
{
|
||||
std::size_t begin = 0;
|
||||
std::size_t nslots = 0;
|
||||
const std::uint8_t * pending = windows_[r].pending_out(begin, nslots);
|
||||
if (nslots == 0)
|
||||
return;
|
||||
round_batch_hdr hdr{};
|
||||
hdr.round = r;
|
||||
hdr.begin = static_cast<std::uint32_t>(begin);
|
||||
hdr.count = static_cast<std::uint32_t>(nslots);
|
||||
const std::size_t body = nslots * slot_bytes_[r];
|
||||
const std::size_t old = payload.size();
|
||||
payload.resize(old + sizeof(hdr) + body);
|
||||
std::memcpy(payload.data() + old, &hdr, sizeof(hdr));
|
||||
if (body != 0)
|
||||
std::memcpy(payload.data() + old + sizeof(hdr), pending, body);
|
||||
windows_[r].mark_flushed(nslots);
|
||||
}
|
||||
|
||||
void exchange_payload(const std::vector<std::uint8_t> & payload)
|
||||
{
|
||||
// Always exchange so a party with nothing new still receives.
|
||||
std::vector<std::uint8_t> theirs;
|
||||
if (self_id_ < peer_id_)
|
||||
{
|
||||
link_.send_bytes(msg::round_batch, payload);
|
||||
theirs = link_.recv_bytes(msg::round_batch);
|
||||
}
|
||||
else if (self_id_ > peer_id_)
|
||||
{
|
||||
theirs = link_.recv_bytes(msg::round_batch);
|
||||
link_.send_bytes(msg::round_batch, payload);
|
||||
}
|
||||
else
|
||||
throw std::invalid_argument("mux_sink flush with self");
|
||||
++link_exchanges_;
|
||||
ingest(theirs);
|
||||
}
|
||||
|
||||
void ingest(const std::vector<std::uint8_t> & bytes)
|
||||
{
|
||||
std::size_t off = 0;
|
||||
while (off < bytes.size())
|
||||
{
|
||||
if (off + sizeof(round_batch_hdr) > bytes.size())
|
||||
throw std::runtime_error("mux_sink truncated header");
|
||||
round_batch_hdr hdr{};
|
||||
std::memcpy(&hdr, bytes.data() + off, sizeof(hdr));
|
||||
off += sizeof(hdr);
|
||||
if (hdr.round >= slot_bytes_.size())
|
||||
throw std::runtime_error("mux_sink bad round");
|
||||
const std::size_t body =
|
||||
static_cast<std::size_t>(hdr.count) * slot_bytes_[hdr.round];
|
||||
if (off + body > bytes.size())
|
||||
throw std::runtime_error("mux_sink truncated body");
|
||||
windows_[hdr.round].accept_peer_at(hdr.begin, bytes.data() + off,
|
||||
hdr.count);
|
||||
off += body;
|
||||
}
|
||||
}
|
||||
|
||||
channel & link_;
|
||||
unsigned self_id_ = 0;
|
||||
unsigned peer_id_ = 0;
|
||||
std::size_t count_ = 0;
|
||||
std::vector<std::size_t> slot_bytes_;
|
||||
std::vector<round_window> windows_;
|
||||
std::uint64_t link_exchanges_ = 0;
|
||||
};
|
||||
|
||||
/// @brief Bind a mux sink to the p0–p1 link of an existing trio.
|
||||
inline mux_sink make_mux_sink(trio & net, std::size_t count,
|
||||
std::vector<std::size_t> slot_bytes)
|
||||
{
|
||||
const role self = net.self();
|
||||
if (self == role::p2)
|
||||
throw std::logic_error("mux_sink is for computing parties");
|
||||
const role peer = self == role::p0 ? role::p1 : role::p0;
|
||||
return mux_sink(net.to(peer), to_u(self), to_u(peer), count,
|
||||
std::move(slot_bytes));
|
||||
}
|
||||
|
||||
/// @brief Default `comm_hook`: framed mesh channels and a mux batch sink.
|
||||
/// @details Subclass to reroute selected calls; leave the rest to these
|
||||
/// defaults. Installing this hook is equivalent to a null hook for
|
||||
/// framing, but lets overrides replace individual methods.
|
||||
class mesh_comm_hook : public comm_hook
|
||||
{
|
||||
public:
|
||||
void send_bytes(trio & net, unsigned peer_id, msg tag, const void * data,
|
||||
std::size_t n) override
|
||||
{
|
||||
net.to(static_cast<role>(peer_id)).send_bytes(tag, data, n);
|
||||
}
|
||||
|
||||
std::vector<std::uint8_t> recv_bytes(trio & net, unsigned peer_id,
|
||||
msg tag) override
|
||||
{
|
||||
return net.to(static_cast<role>(peer_id)).recv_bytes(tag);
|
||||
}
|
||||
|
||||
std::vector<std::uint8_t> exchange_bytes(trio & net, unsigned peer_id,
|
||||
msg tag, const void * data, std::size_t n) override
|
||||
{
|
||||
auto & link = net.to(static_cast<role>(peer_id));
|
||||
const unsigned self_id = to_u(net.self());
|
||||
std::vector<std::uint8_t> theirs(n);
|
||||
if (self_id < peer_id)
|
||||
{
|
||||
link.send_bytes(tag, data, n);
|
||||
theirs = link.recv_bytes(tag);
|
||||
}
|
||||
else if (self_id > peer_id)
|
||||
{
|
||||
theirs = link.recv_bytes(tag);
|
||||
link.send_bytes(tag, data, n);
|
||||
}
|
||||
else
|
||||
throw std::invalid_argument("mesh_comm_hook exchange with self");
|
||||
if (theirs.size() != n)
|
||||
throw std::runtime_error("mesh_comm_hook exchange size");
|
||||
return theirs;
|
||||
}
|
||||
|
||||
std::vector<std::uint8_t> exchange_vec_bytes(trio & net, unsigned peer_id,
|
||||
msg tag, const void * data, std::size_t nbytes) override
|
||||
{
|
||||
return exchange_bytes(net, peer_id, tag, data, nbytes);
|
||||
}
|
||||
|
||||
std::unique_ptr<RoundSink> batch(trio & net, std::size_t count,
|
||||
std::vector<std::size_t> slot_bytes) override
|
||||
{
|
||||
return std::make_unique<mux_sink>(make_mux_sink(net, count,
|
||||
std::move(slot_bytes)));
|
||||
}
|
||||
};
|
||||
|
||||
inline std::unique_ptr<RoundSink> trio::batch(std::size_t count,
|
||||
std::vector<std::size_t> slot_bytes)
|
||||
{
|
||||
if (hook_)
|
||||
return hook_->batch(*this, count, std::move(slot_bytes));
|
||||
return std::make_unique<mux_sink>(make_mux_sink(*this, count,
|
||||
std::move(slot_bytes)));
|
||||
}
|
||||
|
||||
} // namespace net
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_NET_MUX_SINK_HPP__
|
||||
1171
include/dpf/net/party_session.hpp
Normal file
1171
include/dpf/net/party_session.hpp
Normal file
File diff suppressed because it is too large
Load diff
217
include/dpf/net/party_tape_io.hpp
Normal file
217
include/dpf/net/party_tape_io.hpp
Normal file
|
|
@ -0,0 +1,217 @@
|
|||
/// @file dpf/net/party_tape_io.hpp
|
||||
/// @brief Deal / accept a `dpf::beavers::party_tape` over a framed channel.
|
||||
#ifndef LIBDPF_INCLUDE_DPF_NET_PARTY_TAPE_IO_HPP__
|
||||
#define LIBDPF_INCLUDE_DPF_NET_PARTY_TAPE_IO_HPP__
|
||||
|
||||
#include <cstdint>
|
||||
#include <stdexcept>
|
||||
#include <type_traits>
|
||||
#include <vector>
|
||||
|
||||
#include "dpf/beaver.hpp"
|
||||
#include "dpf/net/channel.hpp"
|
||||
#include "dpf/net/trio.hpp"
|
||||
|
||||
namespace dpf
|
||||
{
|
||||
namespace net
|
||||
{
|
||||
namespace detail
|
||||
{
|
||||
|
||||
template <typename Ring>
|
||||
void send_ring_vec(channel & c, const std::vector<Ring> & v)
|
||||
{
|
||||
static_assert(std::is_trivially_copyable_v<Ring>,
|
||||
"party_tape ring must be trivially copyable");
|
||||
std::uint64_t n = v.size();
|
||||
c.send(msg::ring_vector, n);
|
||||
if (n != 0)
|
||||
c.send_vec(v.data(), v.size(), msg::ring_vector);
|
||||
}
|
||||
|
||||
template <typename Ring>
|
||||
std::vector<Ring> recv_ring_vec(channel & c)
|
||||
{
|
||||
auto n = c.recv<std::uint64_t>(msg::ring_vector);
|
||||
if (n == 0)
|
||||
return {};
|
||||
return c.recv_vec<Ring>(msg::ring_vector);
|
||||
}
|
||||
|
||||
inline void send_flags(channel & c, const std::vector<std::uint8_t> & v)
|
||||
{
|
||||
std::uint64_t n = v.size();
|
||||
c.send(msg::bytes, n);
|
||||
if (n != 0)
|
||||
c.send_bytes(msg::bytes, v);
|
||||
}
|
||||
|
||||
inline std::vector<std::uint8_t> recv_flags(channel & c)
|
||||
{
|
||||
auto n = c.recv<std::uint64_t>(msg::bytes);
|
||||
if (n == 0)
|
||||
return {};
|
||||
auto body = c.recv_bytes(msg::bytes);
|
||||
if (body.size() != n)
|
||||
throw std::runtime_error("party_tape flag size mismatch");
|
||||
return body;
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
template <typename Ring>
|
||||
void send_party_tape(channel & c, const beavers::party_tape<Ring> & tape)
|
||||
{
|
||||
c.send(msg::beaver_tape, std::uint8_t{2});
|
||||
c.send(msg::beaver_tape, static_cast<std::uint8_t>(tape.has_mac ? 1 : 0));
|
||||
detail::send_ring_vec(c, tape.lambda);
|
||||
detail::send_flags(c, tape.lambda_ready);
|
||||
detail::send_ring_vec(c, tape.monomial);
|
||||
detail::send_flags(c, tape.monomial_ready);
|
||||
detail::send_ring_vec(c, tape.bundles);
|
||||
detail::send_flags(c, tape.bundles_ready);
|
||||
detail::send_ring_vec(c, tape.dot_cross);
|
||||
detail::send_flags(c, tape.dot_ready);
|
||||
if (tape.has_mac)
|
||||
{
|
||||
detail::send_ring_vec(c, tape.lambda_tag);
|
||||
detail::send_ring_vec(c, tape.monomial_tag);
|
||||
detail::send_ring_vec(c, tape.bundles_tag);
|
||||
detail::send_ring_vec(c, tape.dot_cross_tag);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename Ring>
|
||||
beavers::party_tape<Ring> recv_party_tape(channel & c)
|
||||
{
|
||||
const auto ver = c.recv<std::uint8_t>(msg::beaver_tape);
|
||||
beavers::party_tape<Ring> tape;
|
||||
if (ver >= 2)
|
||||
tape.has_mac = c.recv<std::uint8_t>(msg::beaver_tape) != 0;
|
||||
tape.lambda = detail::recv_ring_vec<Ring>(c);
|
||||
tape.lambda_ready = detail::recv_flags(c);
|
||||
tape.monomial = detail::recv_ring_vec<Ring>(c);
|
||||
tape.monomial_ready = detail::recv_flags(c);
|
||||
tape.bundles = detail::recv_ring_vec<Ring>(c);
|
||||
tape.bundles_ready = detail::recv_flags(c);
|
||||
tape.dot_cross = detail::recv_ring_vec<Ring>(c);
|
||||
tape.dot_ready = detail::recv_flags(c);
|
||||
if (tape.has_mac)
|
||||
{
|
||||
tape.lambda_tag = detail::recv_ring_vec<Ring>(c);
|
||||
tape.monomial_tag = detail::recv_ring_vec<Ring>(c);
|
||||
tape.bundles_tag = detail::recv_ring_vec<Ring>(c);
|
||||
tape.dot_cross_tag = detail::recv_ring_vec<Ring>(c);
|
||||
}
|
||||
return tape;
|
||||
}
|
||||
|
||||
/// @brief Dealer exports and sends each party's tape.
|
||||
/// @tparam Ring Beaver ring
|
||||
/// @param net the connected trio
|
||||
/// @param s the sampled session
|
||||
/// @throws std::logic_error if this process is not p2
|
||||
template <typename Ring>
|
||||
void send_party_tape(trio & net, role peer, const beavers::party_tape<Ring> & tape)
|
||||
{
|
||||
net.send_to(peer, msg::beaver_tape, std::uint8_t{2});
|
||||
net.send_to(peer, msg::beaver_tape,
|
||||
static_cast<std::uint8_t>(tape.has_mac ? 1 : 0));
|
||||
auto send_ring_vec = [&](const std::vector<Ring> & v) {
|
||||
const std::uint64_t n = v.size();
|
||||
net.send_to(peer, msg::ring_vector, n);
|
||||
if (n != 0)
|
||||
net.send_vec_to(peer, msg::ring_vector, v);
|
||||
};
|
||||
auto send_flags = [&](const std::vector<std::uint8_t> & v) {
|
||||
const std::uint64_t n = v.size();
|
||||
net.send_to(peer, msg::bytes, n);
|
||||
if (n != 0)
|
||||
net.send_bytes_to(peer, msg::bytes, v.data(), v.size());
|
||||
};
|
||||
send_ring_vec(tape.lambda);
|
||||
send_flags(tape.lambda_ready);
|
||||
send_ring_vec(tape.monomial);
|
||||
send_flags(tape.monomial_ready);
|
||||
send_ring_vec(tape.bundles);
|
||||
send_flags(tape.bundles_ready);
|
||||
send_ring_vec(tape.dot_cross);
|
||||
send_flags(tape.dot_ready);
|
||||
if (tape.has_mac)
|
||||
{
|
||||
send_ring_vec(tape.lambda_tag);
|
||||
send_ring_vec(tape.monomial_tag);
|
||||
send_ring_vec(tape.bundles_tag);
|
||||
send_ring_vec(tape.dot_cross_tag);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename Ring>
|
||||
beavers::party_tape<Ring> recv_party_tape(trio & net, role peer)
|
||||
{
|
||||
const auto ver = net.recv_from<std::uint8_t>(peer, msg::beaver_tape);
|
||||
beavers::party_tape<Ring> tape;
|
||||
if (ver >= 2)
|
||||
tape.has_mac = net.recv_from<std::uint8_t>(peer, msg::beaver_tape) != 0;
|
||||
auto recv_ring_vec = [&]() {
|
||||
const auto n = net.recv_from<std::uint64_t>(peer, msg::ring_vector);
|
||||
if (n == 0)
|
||||
return std::vector<Ring>{};
|
||||
return net.recv_vec_from<Ring>(peer, msg::ring_vector);
|
||||
};
|
||||
auto recv_flags = [&]() {
|
||||
const auto n = net.recv_from<std::uint64_t>(peer, msg::bytes);
|
||||
if (n == 0)
|
||||
return std::vector<std::uint8_t>{};
|
||||
auto body = net.recv_bytes_from(peer, msg::bytes);
|
||||
if (body.size() != n)
|
||||
throw std::runtime_error("party_tape flag size mismatch");
|
||||
return body;
|
||||
};
|
||||
tape.lambda = recv_ring_vec();
|
||||
tape.lambda_ready = recv_flags();
|
||||
tape.monomial = recv_ring_vec();
|
||||
tape.monomial_ready = recv_flags();
|
||||
tape.bundles = recv_ring_vec();
|
||||
tape.bundles_ready = recv_flags();
|
||||
tape.dot_cross = recv_ring_vec();
|
||||
tape.dot_ready = recv_flags();
|
||||
if (tape.has_mac)
|
||||
{
|
||||
tape.lambda_tag = recv_ring_vec();
|
||||
tape.monomial_tag = recv_ring_vec();
|
||||
tape.bundles_tag = recv_ring_vec();
|
||||
tape.dot_cross_tag = recv_ring_vec();
|
||||
}
|
||||
return tape;
|
||||
}
|
||||
|
||||
template <typename Ring>
|
||||
void deal_session(trio & net, const beavers::session<Ring> & s)
|
||||
{
|
||||
if (net.self() != role::p2)
|
||||
throw std::logic_error("deal_session is for the dealer");
|
||||
send_party_tape(net, role::p0, s.export_party(0));
|
||||
send_party_tape(net, role::p1, s.export_party(1));
|
||||
}
|
||||
|
||||
/// @brief Computing party receives its tape from the dealer.
|
||||
/// @tparam Ring Beaver ring
|
||||
/// @param net the connected trio
|
||||
/// @return this party's tape
|
||||
/// @throws std::logic_error if this process is p2
|
||||
/// @throws std::runtime_error if a flag vector's length disagrees with its header
|
||||
template <typename Ring>
|
||||
HEDLEY_WARN_UNUSED_RESULT
|
||||
beavers::party_tape<Ring> accept_session(trio & net)
|
||||
{
|
||||
if (net.self() == role::p2)
|
||||
throw std::logic_error("dealer does not accept_session");
|
||||
return recv_party_tape<Ring>(net, role::p2);
|
||||
}
|
||||
|
||||
} // namespace net
|
||||
} // namespace dpf
|
||||
|
||||
#endif // LIBDPF_INCLUDE_DPF_NET_PARTY_TAPE_IO_HPP__
|
||||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Add a link
Reference in a new issue