libdpf/include/dpf/cost_pass.hpp
Ryan Henry 0d22946a0e Checkpoint the party/runtime stack before share-program and malicious-mode work.
Ship the TLS mesh, composer, Beaver/Yao/leaf MPC, prep/online paths, apps, and docs so the tree is pushable before elevating share_expr, security_mode, and prep resume.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-09-28 05:59:19 -06:00

155 lines
4.2 KiB
C++

/// @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__