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
633
party/oblivious_hash.hpp
Normal file
633
party/oblivious_hash.hpp
Normal file
|
|
@ -0,0 +1,633 @@
|
|||
/// @file party/oblivious_hash.hpp
|
||||
/// @brief Correction-seed hash that keeps the path prefix shared.
|
||||
/// @details Matches `detail::vdpf::hash_node` without either party learning
|
||||
/// the prefix. The opened value is only `H(s0) XOR H(s1)`.
|
||||
/// SubBytes uses the Boyar–Peralta AES S-box (ePrint 2011/332;
|
||||
/// Yale CMT SLP with 32 ANDs). Bit×bit ANDs use bit Beaver triples
|
||||
/// and are scheduled by multiplicative depth (six exchanges per
|
||||
/// SubBytes), so one level costs `8×16×10×32 = 40960` bit-ANDs.
|
||||
|
||||
#ifndef LIBDPF_PARTY_OBLIVIOUS_HASH_HPP__
|
||||
#define LIBDPF_PARTY_OBLIVIOUS_HASH_HPP__
|
||||
|
||||
#include <array>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
#include <stdexcept>
|
||||
#include <vector>
|
||||
|
||||
#include "simde/simde/x86/avx2.h"
|
||||
|
||||
#include "aes_mmo_ref.hpp"
|
||||
#include "aes_sbox_bp.hpp"
|
||||
#include "dpf/doerner_shelat.hpp"
|
||||
#include "dpf/prg_aes.hpp"
|
||||
#include "dpf/verifiable.hpp"
|
||||
#include "dpf/net/round_sink.hpp"
|
||||
#include "dpf/net/sink_exchange.hpp"
|
||||
|
||||
/* aes_bp is file-scope */
|
||||
|
||||
/// @brief AND-triples consumed by one level of `oblivious_cs`.
|
||||
inline constexpr std::size_t hash_level_and_count() noexcept
|
||||
{
|
||||
// Eight AES (two seeds × four MMO lanes) × 16 bytes × 10 SubBytes
|
||||
// × 32 Boyar–Peralta ANDs.
|
||||
return 8u * 16u * 10u * aes_bp::and_count;
|
||||
}
|
||||
|
||||
inline std::vector<bit_and_pad_msg> sample_and_tape(std::size_t n, int party)
|
||||
{
|
||||
::dpf::detail::urandom_pad_rng pads;
|
||||
std::vector<bit_and_pad_msg> out(n);
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
const auto full = ::dpf::detail::ds_sample_bit_and(pads);
|
||||
out[i] = bit_pad_for(full, party);
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
inline std::vector<std::uint8_t> pack_mask_bits(
|
||||
const std::uint8_t * d, const std::uint8_t * e, std::size_t n)
|
||||
{
|
||||
std::vector<std::uint8_t> out((2u * n + 7u) / 8u);
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
const std::size_t db = 2u * i;
|
||||
const std::size_t eb = db + 1u;
|
||||
if (d[i] & 1u)
|
||||
out[db / 8u] = static_cast<std::uint8_t>(
|
||||
out[db / 8u] | static_cast<std::uint8_t>(1u << (db % 8u)));
|
||||
if (e[i] & 1u)
|
||||
out[eb / 8u] = static_cast<std::uint8_t>(
|
||||
out[eb / 8u] | static_cast<std::uint8_t>(1u << (eb % 8u)));
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
inline void unpack_mask_bits(const std::vector<std::uint8_t> & packed,
|
||||
std::uint8_t * d, std::uint8_t * e, std::size_t n)
|
||||
{
|
||||
if (packed.size() != (2u * n + 7u) / 8u)
|
||||
throw std::runtime_error("oblivious AND packed size");
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
const std::size_t db = 2u * i;
|
||||
const std::size_t eb = db + 1u;
|
||||
d[i] = static_cast<std::uint8_t>(
|
||||
(packed[db / 8u] >> (db % 8u)) & 1u);
|
||||
e[i] = static_cast<std::uint8_t>(
|
||||
(packed[eb / 8u] >> (eb % 8u)) & 1u);
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief One exchange of packed `(x⊕α, y⊕β)` masks for `n` bit-ANDs.
|
||||
template <int Me>
|
||||
std::vector<std::uint8_t> batch_bit_and(net::trio & net,
|
||||
const std::uint8_t * x, const std::uint8_t * y, std::size_t n,
|
||||
const bit_and_pad_msg * triples)
|
||||
{
|
||||
const role peer = Me == 0 ? role::p1 : role::p0;
|
||||
std::vector<std::uint8_t> d(n), e(n);
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
d[i] = static_cast<std::uint8_t>(x[i] ^ triples[i].a);
|
||||
e[i] = static_cast<std::uint8_t>(y[i] ^ triples[i].b);
|
||||
}
|
||||
const auto mine = pack_mask_bits(d.data(), e.data(), n);
|
||||
std::vector<std::uint8_t> peer_bytes;
|
||||
if constexpr (Me == 0)
|
||||
{
|
||||
net.send_bytes_to(peer, net::msg::delta, mine.data(), mine.size());
|
||||
peer_bytes = net.recv_bytes_from(peer, net::msg::delta);
|
||||
}
|
||||
else
|
||||
{
|
||||
peer_bytes = net.recv_bytes_from(peer, net::msg::delta);
|
||||
net.send_bytes_to(peer, net::msg::delta, mine.data(), mine.size());
|
||||
}
|
||||
std::vector<std::uint8_t> pd(n), pe(n);
|
||||
unpack_mask_bits(peer_bytes, pd.data(), pe.data(), n);
|
||||
|
||||
// Standard Beaver: z = [c] ⊕ (d∧[b]) ⊕ ([a]∧e) ⊕ (d∧e), and the public
|
||||
// `d∧e` term is added on one party only. Adding it on both cancels under
|
||||
// XOR and the S-box (and every correction seed) diverges from `hash_node`.
|
||||
std::vector<std::uint8_t> shares(n);
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
const std::uint8_t od = static_cast<std::uint8_t>(d[i] ^ pd[i]);
|
||||
const std::uint8_t oe = static_cast<std::uint8_t>(e[i] ^ pe[i]);
|
||||
shares[i] = detail::ds_bit_and_party(
|
||||
od, oe, triples[i].a, triples[i].b, triples[i].c, Me == 0);
|
||||
}
|
||||
return shares;
|
||||
}
|
||||
|
||||
/// @brief Packed bit-AND layer on a RoundSink. Slot must fit `4 + packed`.
|
||||
template <int Me>
|
||||
std::vector<std::uint8_t> batch_bit_and_on_sink(dpf::net::RoundSink & sink,
|
||||
std::uint16_t & round, std::size_t index, const std::uint8_t * x,
|
||||
const std::uint8_t * y, std::size_t n, const bit_and_pad_msg * triples)
|
||||
{
|
||||
std::vector<std::uint8_t> d(n), e(n);
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
d[i] = static_cast<std::uint8_t>(x[i] ^ triples[i].a);
|
||||
e[i] = static_cast<std::uint8_t>(y[i] ^ triples[i].b);
|
||||
}
|
||||
const auto mine = pack_mask_bits(d.data(), e.data(), n);
|
||||
if (round >= sink.rounds())
|
||||
throw std::runtime_error("batch_bit_and_on_sink: out of rounds");
|
||||
const std::size_t slot = sink.slot_bytes(round);
|
||||
const std::size_t need = sizeof(std::uint32_t) + mine.size();
|
||||
if (need > slot)
|
||||
throw std::runtime_error("batch_bit_and_on_sink: slot too small");
|
||||
std::vector<std::uint8_t> buf(slot, 0);
|
||||
const std::uint32_t nbytes = static_cast<std::uint32_t>(mine.size());
|
||||
std::memcpy(buf.data(), &nbytes, 4);
|
||||
if (!mine.empty())
|
||||
std::memcpy(buf.data() + 4, mine.data(), mine.size());
|
||||
sink.submit(round, index, buf.data(), buf.size());
|
||||
sink.flush_round(round);
|
||||
sink.poll();
|
||||
if (!sink.peer_ready(round, index))
|
||||
throw std::runtime_error("batch_bit_and_on_sink: peer missing");
|
||||
std::vector<std::uint8_t> peer_buf(slot);
|
||||
sink.read_peer(round, index, peer_buf.data(), slot);
|
||||
std::uint32_t pn = 0;
|
||||
std::memcpy(&pn, peer_buf.data(), 4);
|
||||
if (pn != nbytes)
|
||||
throw std::runtime_error("batch_bit_and_on_sink: peer size");
|
||||
std::vector<std::uint8_t> peer_bytes(pn);
|
||||
if (pn != 0)
|
||||
std::memcpy(peer_bytes.data(), peer_buf.data() + 4, pn);
|
||||
++round;
|
||||
std::vector<std::uint8_t> pd(n), pe(n);
|
||||
unpack_mask_bits(peer_bytes, pd.data(), pe.data(), n);
|
||||
std::vector<std::uint8_t> shares(n);
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
{
|
||||
const std::uint8_t od = static_cast<std::uint8_t>(d[i] ^ pd[i]);
|
||||
const std::uint8_t oe = static_cast<std::uint8_t>(e[i] ^ pe[i]);
|
||||
shares[i] = detail::ds_bit_and_party(
|
||||
od, oe, triples[i].a, triples[i].b, triples[i].c, Me == 0);
|
||||
}
|
||||
return shares;
|
||||
}
|
||||
|
||||
/// @brief Shared SubBytes via Boyar–Peralta (32 ANDs per byte, depth-batched).
|
||||
template <int Me>
|
||||
void sub_bytes_shared(net::trio & net, std::uint8_t * state, std::size_t bytes,
|
||||
const bit_and_pad_msg * & tape)
|
||||
{
|
||||
std::vector<std::uint8_t> wire(bytes * aes_bp::wire_count);
|
||||
auto at = [&](std::size_t b, std::size_t w) -> std::uint8_t & {
|
||||
return wire[b * aes_bp::wire_count + w];
|
||||
};
|
||||
for (std::size_t b = 0; b < bytes; ++b)
|
||||
{
|
||||
for (int i = 0; i < 8; ++i)
|
||||
at(b, static_cast<std::size_t>(i)) =
|
||||
static_cast<std::uint8_t>((state[b] >> (7 - i)) & 1u);
|
||||
}
|
||||
|
||||
std::array<bool, aes_bp::wire_count> ready{};
|
||||
for (std::size_t i = 0; i < 8; ++i)
|
||||
ready[i] = true;
|
||||
std::array<bool, aes_bp::op_count> done{};
|
||||
std::size_t finished = 0;
|
||||
while (finished < aes_bp::op_count)
|
||||
{
|
||||
bool progress = true;
|
||||
while (progress)
|
||||
{
|
||||
progress = false;
|
||||
for (std::size_t oi = 0; oi < aes_bp::op_count; ++oi)
|
||||
{
|
||||
if (done[oi])
|
||||
continue;
|
||||
const auto kind = aes_bp::ops[oi][0];
|
||||
if (kind == 1)
|
||||
continue;
|
||||
const auto dst = aes_bp::ops[oi][1];
|
||||
const auto a = aes_bp::ops[oi][2];
|
||||
const auto bb = aes_bp::ops[oi][3];
|
||||
if (!ready[a] || !ready[bb])
|
||||
continue;
|
||||
for (std::size_t b = 0; b < bytes; ++b)
|
||||
{
|
||||
auto v = static_cast<std::uint8_t>(at(b, a) ^ at(b, bb));
|
||||
if (kind == 2)
|
||||
{
|
||||
if constexpr (Me == 0)
|
||||
v = static_cast<std::uint8_t>(v ^ 1u);
|
||||
}
|
||||
at(b, dst) = v;
|
||||
}
|
||||
ready[dst] = true;
|
||||
done[oi] = true;
|
||||
++finished;
|
||||
progress = true;
|
||||
}
|
||||
}
|
||||
|
||||
std::vector<std::size_t> ands;
|
||||
ands.reserve(aes_bp::and_count);
|
||||
for (std::size_t oi = 0; oi < aes_bp::op_count; ++oi)
|
||||
{
|
||||
if (done[oi] || aes_bp::ops[oi][0] != 1)
|
||||
continue;
|
||||
const auto a = aes_bp::ops[oi][2];
|
||||
const auto bb = aes_bp::ops[oi][3];
|
||||
if (ready[a] && ready[bb])
|
||||
ands.push_back(oi);
|
||||
}
|
||||
if (ands.empty())
|
||||
{
|
||||
if (finished != aes_bp::op_count)
|
||||
throw std::runtime_error("oblivious SubBytes stuck");
|
||||
break;
|
||||
}
|
||||
|
||||
const std::size_t n = ands.size() * bytes;
|
||||
std::vector<std::uint8_t> xs(n), ys(n);
|
||||
for (std::size_t ai = 0; ai < ands.size(); ++ai)
|
||||
{
|
||||
const auto a = aes_bp::ops[ands[ai]][2];
|
||||
const auto bb = aes_bp::ops[ands[ai]][3];
|
||||
for (std::size_t b = 0; b < bytes; ++b)
|
||||
{
|
||||
xs[ai * bytes + b] = at(b, a);
|
||||
ys[ai * bytes + b] = at(b, bb);
|
||||
}
|
||||
}
|
||||
const auto prod = batch_bit_and<Me>(net, xs.data(), ys.data(), n, tape);
|
||||
tape += n;
|
||||
for (std::size_t ai = 0; ai < ands.size(); ++ai)
|
||||
{
|
||||
const auto dst = aes_bp::ops[ands[ai]][1];
|
||||
for (std::size_t b = 0; b < bytes; ++b)
|
||||
at(b, dst) = static_cast<std::uint8_t>(prod[ai * bytes + b] & 1u);
|
||||
ready[dst] = true;
|
||||
done[ands[ai]] = true;
|
||||
++finished;
|
||||
}
|
||||
}
|
||||
|
||||
for (std::size_t b = 0; b < bytes; ++b)
|
||||
{
|
||||
std::uint8_t out = 0;
|
||||
for (int i = 0; i < 8; ++i)
|
||||
out = static_cast<std::uint8_t>(
|
||||
out | (at(b, aes_bp::out_wire[static_cast<std::size_t>(i)])
|
||||
<< (7 - i)));
|
||||
state[b] = out;
|
||||
}
|
||||
}
|
||||
|
||||
/// @brief SubBytes with AND layers on a RoundSink (multi-instance friendly).
|
||||
template <int Me>
|
||||
void sub_bytes_shared_on_sink(dpf::net::RoundSink & sink, std::uint16_t & round,
|
||||
std::size_t index, std::uint8_t * state, std::size_t bytes,
|
||||
const bit_and_pad_msg * & tape)
|
||||
{
|
||||
std::vector<std::uint8_t> wire(bytes * aes_bp::wire_count);
|
||||
auto at = [&](std::size_t b, std::size_t w) -> std::uint8_t & {
|
||||
return wire[b * aes_bp::wire_count + w];
|
||||
};
|
||||
for (std::size_t b = 0; b < bytes; ++b)
|
||||
{
|
||||
for (int i = 0; i < 8; ++i)
|
||||
at(b, static_cast<std::size_t>(i)) =
|
||||
static_cast<std::uint8_t>((state[b] >> (7 - i)) & 1u);
|
||||
}
|
||||
|
||||
std::array<bool, aes_bp::wire_count> ready{};
|
||||
for (std::size_t i = 0; i < 8; ++i)
|
||||
ready[i] = true;
|
||||
std::array<bool, aes_bp::op_count> done{};
|
||||
std::size_t finished = 0;
|
||||
while (finished < aes_bp::op_count)
|
||||
{
|
||||
bool progress = true;
|
||||
while (progress)
|
||||
{
|
||||
progress = false;
|
||||
for (std::size_t oi = 0; oi < aes_bp::op_count; ++oi)
|
||||
{
|
||||
if (done[oi])
|
||||
continue;
|
||||
const auto kind = aes_bp::ops[oi][0];
|
||||
if (kind == 1)
|
||||
continue;
|
||||
const auto dst = aes_bp::ops[oi][1];
|
||||
const auto a = aes_bp::ops[oi][2];
|
||||
const auto bb = aes_bp::ops[oi][3];
|
||||
if (!ready[a] || !ready[bb])
|
||||
continue;
|
||||
for (std::size_t b = 0; b < bytes; ++b)
|
||||
{
|
||||
auto v = static_cast<std::uint8_t>(at(b, a) ^ at(b, bb));
|
||||
if (kind == 2)
|
||||
{
|
||||
if constexpr (Me == 0)
|
||||
v = static_cast<std::uint8_t>(v ^ 1u);
|
||||
}
|
||||
at(b, dst) = v;
|
||||
}
|
||||
ready[dst] = true;
|
||||
done[oi] = true;
|
||||
++finished;
|
||||
progress = true;
|
||||
}
|
||||
}
|
||||
|
||||
std::vector<std::size_t> ands;
|
||||
ands.reserve(aes_bp::and_count);
|
||||
for (std::size_t oi = 0; oi < aes_bp::op_count; ++oi)
|
||||
{
|
||||
if (done[oi] || aes_bp::ops[oi][0] != 1)
|
||||
continue;
|
||||
const auto a = aes_bp::ops[oi][2];
|
||||
const auto bb = aes_bp::ops[oi][3];
|
||||
if (ready[a] && ready[bb])
|
||||
ands.push_back(oi);
|
||||
}
|
||||
if (ands.empty())
|
||||
{
|
||||
if (finished != aes_bp::op_count)
|
||||
throw std::runtime_error("oblivious SubBytes stuck");
|
||||
break;
|
||||
}
|
||||
|
||||
const std::size_t n = ands.size() * bytes;
|
||||
std::vector<std::uint8_t> xs(n), ys(n);
|
||||
for (std::size_t ai = 0; ai < ands.size(); ++ai)
|
||||
{
|
||||
const auto a = aes_bp::ops[ands[ai]][2];
|
||||
const auto bb = aes_bp::ops[ands[ai]][3];
|
||||
for (std::size_t b = 0; b < bytes; ++b)
|
||||
{
|
||||
xs[ai * bytes + b] = at(b, a);
|
||||
ys[ai * bytes + b] = at(b, bb);
|
||||
}
|
||||
}
|
||||
const auto prod = batch_bit_and_on_sink<Me>(
|
||||
sink, round, index, xs.data(), ys.data(), n, tape);
|
||||
tape += n;
|
||||
for (std::size_t ai = 0; ai < ands.size(); ++ai)
|
||||
{
|
||||
const auto dst = aes_bp::ops[ands[ai]][1];
|
||||
for (std::size_t b = 0; b < bytes; ++b)
|
||||
at(b, dst) = static_cast<std::uint8_t>(prod[ai * bytes + b] & 1u);
|
||||
ready[dst] = true;
|
||||
done[ands[ai]] = true;
|
||||
++finished;
|
||||
}
|
||||
}
|
||||
|
||||
for (std::size_t b = 0; b < bytes; ++b)
|
||||
{
|
||||
std::uint8_t out = 0;
|
||||
for (int i = 0; i < 8; ++i)
|
||||
out = static_cast<std::uint8_t>(
|
||||
out | (at(b, aes_bp::out_wire[static_cast<std::size_t>(i)])
|
||||
<< (7 - i)));
|
||||
state[b] = out;
|
||||
}
|
||||
}
|
||||
|
||||
inline void shift_rows_bytes(std::uint8_t s[16])
|
||||
{
|
||||
::dpf::party::aes_ref::shift_rows(s);
|
||||
}
|
||||
|
||||
inline void mix_columns_bytes(std::uint8_t s[16])
|
||||
{
|
||||
::dpf::party::aes_ref::mix_columns(s);
|
||||
}
|
||||
|
||||
inline void xor_bytes(std::uint8_t * dst, const std::uint8_t * src, std::size_t n)
|
||||
{
|
||||
for (std::size_t i = 0; i < n; ++i)
|
||||
dst[i] = static_cast<std::uint8_t>(dst[i] ^ src[i]);
|
||||
}
|
||||
|
||||
inline std::array<std::uint8_t, 16> m128_bytes(simde__m128i v)
|
||||
{
|
||||
std::array<std::uint8_t, 16> b{};
|
||||
std::memcpy(b.data(), &v, 16);
|
||||
return b;
|
||||
}
|
||||
|
||||
/// @brief Public correction seed for one level. Prefix shares stay local.
|
||||
template <int Me>
|
||||
cs_block oblivious_cs(net::trio & net, std::size_t level,
|
||||
std::uint64_t prefix_share, simde__m128i my_seed,
|
||||
const bit_and_pad_msg * tape)
|
||||
{
|
||||
const auto schedule = prg::aes128_key(simde_mm_setzero_si128());
|
||||
const auto rk = [&](int r) {
|
||||
return m128_bytes(schedule.rd_key[static_cast<std::size_t>(r)]);
|
||||
};
|
||||
|
||||
constexpr std::size_t blocks = 8;
|
||||
std::uint8_t state[blocks][16]{};
|
||||
std::uint8_t feed[blocks][16]{};
|
||||
|
||||
auto place_prefix = [&](std::uint8_t * dst) {
|
||||
for (int i = 0; i < 8; ++i)
|
||||
dst[i] = static_cast<std::uint8_t>(
|
||||
dst[i] ^ static_cast<std::uint8_t>(prefix_share >> (8 * i)));
|
||||
};
|
||||
|
||||
for (int owner = 0; owner < 2; ++owner)
|
||||
{
|
||||
const bool mine = (owner == Me);
|
||||
for (int pos = 0; pos < 4; ++pos)
|
||||
{
|
||||
const std::size_t bi = static_cast<std::size_t>(owner * 4 + pos);
|
||||
if (mine)
|
||||
{
|
||||
auto bytes = m128_bytes(my_seed);
|
||||
const std::uint64_t tag = 0x5600ull | (level & 0xffffu);
|
||||
for (int i = 0; i < 8; ++i)
|
||||
bytes[8 + i] = static_cast<std::uint8_t>(
|
||||
bytes[8 + i]
|
||||
^ static_cast<std::uint8_t>(tag >> (8 * i)));
|
||||
place_prefix(bytes.data());
|
||||
std::memcpy(feed[bi], bytes.data(), 16);
|
||||
std::memcpy(state[bi], bytes.data(), 16);
|
||||
}
|
||||
else
|
||||
{
|
||||
std::uint8_t bytes[16]{};
|
||||
place_prefix(bytes);
|
||||
std::memcpy(feed[bi], bytes, 16);
|
||||
std::memcpy(state[bi], bytes, 16);
|
||||
}
|
||||
if constexpr (Me == 0)
|
||||
{
|
||||
auto key0 = rk(0);
|
||||
key0[0] = static_cast<std::uint8_t>(
|
||||
key0[0] ^ static_cast<std::uint8_t>(pos));
|
||||
key0[1] = static_cast<std::uint8_t>(
|
||||
key0[1]
|
||||
^ static_cast<std::uint8_t>(static_cast<unsigned>(pos) >> 8));
|
||||
key0[2] = static_cast<std::uint8_t>(
|
||||
key0[2]
|
||||
^ static_cast<std::uint8_t>(
|
||||
static_cast<unsigned>(pos) >> 16));
|
||||
key0[3] = static_cast<std::uint8_t>(
|
||||
key0[3]
|
||||
^ static_cast<std::uint8_t>(
|
||||
static_cast<unsigned>(pos) >> 24));
|
||||
xor_bytes(state[bi], key0.data(), 16);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
auto all_bytes = [&]() -> std::uint8_t * { return &state[0][0]; };
|
||||
const bit_and_pad_msg * cursor = tape;
|
||||
for (int round = 1; round <= 10; ++round)
|
||||
{
|
||||
sub_bytes_shared<Me>(net, all_bytes(), blocks * 16, cursor);
|
||||
for (std::size_t b = 0; b < blocks; ++b)
|
||||
shift_rows_bytes(state[b]);
|
||||
if (round < 10)
|
||||
{
|
||||
for (std::size_t b = 0; b < blocks; ++b)
|
||||
mix_columns_bytes(state[b]);
|
||||
}
|
||||
if constexpr (Me == 0)
|
||||
{
|
||||
const auto key = rk(round);
|
||||
for (std::size_t b = 0; b < blocks; ++b)
|
||||
xor_bytes(state[b], key.data(), 16);
|
||||
}
|
||||
}
|
||||
if (cursor != tape + hash_level_and_count())
|
||||
throw std::runtime_error("oblivious hash consumed the wrong AND count");
|
||||
|
||||
for (std::size_t b = 0; b < blocks; ++b)
|
||||
xor_bytes(state[b], feed[b], 16);
|
||||
|
||||
cs_block mine{};
|
||||
for (int pos = 0; pos < 4; ++pos)
|
||||
{
|
||||
std::uint8_t mixed[16]{};
|
||||
xor_bytes(mixed, state[static_cast<std::size_t>(pos)], 16);
|
||||
xor_bytes(mixed, state[static_cast<std::size_t>(4 + pos)], 16);
|
||||
std::memcpy(&mine[static_cast<std::size_t>(pos)], mixed, 16);
|
||||
}
|
||||
return open_cs<Me>(net, mine);
|
||||
}
|
||||
|
||||
/// @brief `oblivious_cs` with AND layers on a RoundSink; final open stays on trio.
|
||||
template <int Me>
|
||||
cs_block oblivious_cs_on_sink(net::trio & net, dpf::net::RoundSink & sink,
|
||||
std::uint16_t & round, std::size_t index, std::size_t level,
|
||||
std::uint64_t prefix_share, simde__m128i my_seed,
|
||||
const bit_and_pad_msg * tape)
|
||||
{
|
||||
const auto schedule = prg::aes128_key(simde_mm_setzero_si128());
|
||||
const auto rk = [&](int r) {
|
||||
return m128_bytes(schedule.rd_key[static_cast<std::size_t>(r)]);
|
||||
};
|
||||
|
||||
constexpr std::size_t blocks = 8;
|
||||
std::uint8_t state[blocks][16]{};
|
||||
std::uint8_t feed[blocks][16]{};
|
||||
|
||||
auto place_prefix = [&](std::uint8_t * dst) {
|
||||
for (int i = 0; i < 8; ++i)
|
||||
dst[i] = static_cast<std::uint8_t>(
|
||||
dst[i] ^ static_cast<std::uint8_t>(prefix_share >> (8 * i)));
|
||||
};
|
||||
|
||||
for (int owner = 0; owner < 2; ++owner)
|
||||
{
|
||||
const bool mine = (owner == Me);
|
||||
for (int pos = 0; pos < 4; ++pos)
|
||||
{
|
||||
const std::size_t bi = static_cast<std::size_t>(owner * 4 + pos);
|
||||
if (mine)
|
||||
{
|
||||
auto bytes = m128_bytes(my_seed);
|
||||
const std::uint64_t tag = 0x5600ull | (level & 0xffffu);
|
||||
for (int i = 0; i < 8; ++i)
|
||||
bytes[8 + i] = static_cast<std::uint8_t>(
|
||||
bytes[8 + i]
|
||||
^ static_cast<std::uint8_t>(tag >> (8 * i)));
|
||||
place_prefix(bytes.data());
|
||||
std::memcpy(feed[bi], bytes.data(), 16);
|
||||
std::memcpy(state[bi], bytes.data(), 16);
|
||||
}
|
||||
else
|
||||
{
|
||||
std::uint8_t bytes[16]{};
|
||||
place_prefix(bytes);
|
||||
std::memcpy(feed[bi], bytes, 16);
|
||||
std::memcpy(state[bi], bytes, 16);
|
||||
}
|
||||
if constexpr (Me == 0)
|
||||
{
|
||||
auto key0 = rk(0);
|
||||
key0[0] = static_cast<std::uint8_t>(
|
||||
key0[0] ^ static_cast<std::uint8_t>(pos));
|
||||
key0[1] = static_cast<std::uint8_t>(
|
||||
key0[1]
|
||||
^ static_cast<std::uint8_t>(static_cast<unsigned>(pos) >> 8));
|
||||
key0[2] = static_cast<std::uint8_t>(
|
||||
key0[2]
|
||||
^ static_cast<std::uint8_t>(
|
||||
static_cast<unsigned>(pos) >> 16));
|
||||
key0[3] = static_cast<std::uint8_t>(
|
||||
key0[3]
|
||||
^ static_cast<std::uint8_t>(
|
||||
static_cast<unsigned>(pos) >> 24));
|
||||
xor_bytes(state[bi], key0.data(), 16);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
auto all_bytes = [&]() -> std::uint8_t * { return &state[0][0]; };
|
||||
const bit_and_pad_msg * cursor = tape;
|
||||
for (int aes_round = 1; aes_round <= 10; ++aes_round)
|
||||
{
|
||||
sub_bytes_shared_on_sink<Me>(
|
||||
sink, round, index, all_bytes(), blocks * 16, cursor);
|
||||
for (std::size_t b = 0; b < blocks; ++b)
|
||||
shift_rows_bytes(state[b]);
|
||||
if (aes_round < 10)
|
||||
{
|
||||
for (std::size_t b = 0; b < blocks; ++b)
|
||||
mix_columns_bytes(state[b]);
|
||||
}
|
||||
if constexpr (Me == 0)
|
||||
{
|
||||
const auto key = rk(aes_round);
|
||||
for (std::size_t b = 0; b < blocks; ++b)
|
||||
xor_bytes(state[b], key.data(), 16);
|
||||
}
|
||||
}
|
||||
if (cursor != tape + hash_level_and_count())
|
||||
throw std::runtime_error("oblivious hash consumed the wrong AND count");
|
||||
|
||||
for (std::size_t b = 0; b < blocks; ++b)
|
||||
xor_bytes(state[b], feed[b], 16);
|
||||
|
||||
cs_block mine{};
|
||||
for (int pos = 0; pos < 4; ++pos)
|
||||
{
|
||||
std::uint8_t mixed[16]{};
|
||||
xor_bytes(mixed, state[static_cast<std::size_t>(pos)], 16);
|
||||
xor_bytes(mixed, state[static_cast<std::size_t>(4 + pos)], 16);
|
||||
std::memcpy(&mine[static_cast<std::size_t>(pos)], mixed, 16);
|
||||
}
|
||||
return open_cs<Me>(net, mine);
|
||||
}
|
||||
|
||||
#endif
|
||||
Loading…
Add table
Add a link
Reference in a new issue