libdpf/test/tests/shared_and_test.cpp
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

93 lines
3.7 KiB
C++

#include <gtest/gtest.h>
#include <cstdint>
#include "aes_mmo_ref.hpp"
#include "aes_sbox_bp.hpp"
#include "dpf/doerner_shelat.hpp"
namespace
{
std::uint8_t sbox_circuit(std::uint8_t x)
{
std::uint8_t w[aes_bp::wire_count]{};
for (int i = 0; i < 8; ++i)
w[i] = static_cast<std::uint8_t>((x >> (7 - i)) & 1u);
for (std::size_t oi = 0; oi < aes_bp::op_count; ++oi)
{
const auto kind = aes_bp::ops[oi][0];
const auto dst = aes_bp::ops[oi][1];
const auto a = aes_bp::ops[oi][2];
const auto b = aes_bp::ops[oi][3];
if (kind == 0)
w[dst] = static_cast<std::uint8_t>(w[a] ^ w[b]);
else if (kind == 1)
w[dst] = static_cast<std::uint8_t>(w[a] & w[b]);
else
w[dst] = static_cast<std::uint8_t>(w[a] ^ w[b] ^ 1u);
}
std::uint8_t out = 0;
for (int i = 0; i < 8; ++i)
out = static_cast<std::uint8_t>(
out | (w[aes_bp::out_wire[static_cast<std::size_t>(i)]] << (7 - i)));
return out;
}
} // namespace
TEST(SharedAnd, BoyarPeraltaMatchesAesSbox)
{
int ands = 0;
for (std::size_t oi = 0; oi < aes_bp::op_count; ++oi)
if (aes_bp::ops[oi][0] == 1)
++ands;
EXPECT_EQ(ands, static_cast<int>(aes_bp::and_count));
for (int x = 0; x < 256; ++x)
{
const auto byte = static_cast<std::uint8_t>(x);
EXPECT_EQ(sbox_circuit(byte), dpf::party::aes_ref::sbox_at(byte))
<< "byte " << x;
}
}
TEST(SharedAnd, PublicProductTermIsAddedOnce)
{
// Every XOR-share of a bit AND. The opened masks are `d = x XOR a` and
// `e = y XOR b`. Party 0 alone adds the public `d AND e`. Adding it on
// both parties cancels that term and the product is wrong whenever it
// is 1.
int cancelled = 0;
for (unsigned bits = 0; bits < 256; ++bits)
{
const std::uint8_t a = static_cast<std::uint8_t>(bits & 1u);
const std::uint8_t b = static_cast<std::uint8_t>((bits >> 1) & 1u);
const std::uint8_t x = static_cast<std::uint8_t>((bits >> 2) & 1u);
const std::uint8_t y = static_cast<std::uint8_t>((bits >> 3) & 1u);
const std::uint8_t a0 = static_cast<std::uint8_t>((bits >> 4) & 1u);
const std::uint8_t b0 = static_cast<std::uint8_t>((bits >> 5) & 1u);
const std::uint8_t x0 = static_cast<std::uint8_t>((bits >> 6) & 1u);
const std::uint8_t y0 = static_cast<std::uint8_t>((bits >> 7) & 1u);
const std::uint8_t c = static_cast<std::uint8_t>(a & b);
const std::uint8_t c0 = static_cast<std::uint8_t>((a0 & b0) ^ ((bits * 3u) & 1u));
const std::uint8_t a1 = static_cast<std::uint8_t>(a0 ^ a);
const std::uint8_t b1 = static_cast<std::uint8_t>(b0 ^ b);
const std::uint8_t c1 = static_cast<std::uint8_t>(c0 ^ c);
const std::uint8_t x1 = static_cast<std::uint8_t>(x0 ^ x);
const std::uint8_t y1 = static_cast<std::uint8_t>(y0 ^ y);
const std::uint8_t d = static_cast<std::uint8_t>((x0 ^ a0) ^ (x1 ^ a1));
const std::uint8_t e = static_cast<std::uint8_t>((y0 ^ b0) ^ (y1 ^ b1));
const std::uint8_t z0 = dpf::detail::ds_bit_and_party(d, e, a0, b0, c0, true);
const std::uint8_t z1 = dpf::detail::ds_bit_and_party(d, e, a1, b1, c1, false);
EXPECT_EQ(static_cast<std::uint8_t>(z0 ^ z1), static_cast<std::uint8_t>(x & y));
const std::uint8_t both = static_cast<std::uint8_t>(
dpf::detail::ds_bit_and_party(d, e, a0, b0, c0, true)
^ dpf::detail::ds_bit_and_party(d, e, a1, b1, c1, true));
if ((d & e) != 0)
{
EXPECT_NE(both, static_cast<std::uint8_t>(x & y));
++cancelled;
}
}
EXPECT_GT(cancelled, 0);
}