libdpf/test/tests/yao_stack_test.cpp

315 lines
10 KiB
C++
Raw Normal View History

#include <gtest/gtest.h>
#include <cstdint>
#include <vector>
#include "dpf/yao.hpp"
#include "dpf/yao_stack.hpp"
namespace
{
using dpf::yao::bit;
using dpf::yao::netlist;
netlist and_n(unsigned n)
{
netlist nl;
std::vector<bit> in;
in.reserve(n);
for (unsigned i = 0; i < n; ++i)
in.push_back(nl.shared_in());
bit acc = in[0];
for (unsigned i = 1; i < n; ++i)
acc = nl.and_(acc, in[i]);
nl.out(acc);
return nl;
}
netlist xor2()
{
netlist nl;
auto a = nl.shared_in();
auto b = nl.shared_in();
nl.out(nl.xor_(a, b));
return nl;
}
std::uint8_t plain_one(const netlist & nl, const std::vector<std::uint8_t> & in)
{
return dpf::yao::eval_plain(nl, in.data())[0];
}
} // namespace
TEST(YaoStack, IfCostsTheHeavierBranch)
{
auto heavy = and_n(3);
auto light = and_n(2);
EXPECT_EQ(heavy.n_and(), 2u);
EXPECT_EQ(light.n_and(), 1u);
std::vector<std::uint8_t> h0{1, 1, 1};
std::vector<std::uint8_t> h1{0, 0, 0};
std::vector<std::uint8_t> l0{1, 1};
std::vector<std::uint8_t> l1{0, 1};
for (std::uint8_t c0 = 0; c0 < 2; ++c0)
{
for (std::uint8_t c1 = 0; c1 < 2; ++c1)
{
auto got = dpf::yao::eval_if(heavy, light, c0, c1,
h0.data(), h1.data(), l0.data(), l1.data());
EXPECT_EQ(got.stack_blocks, 4u);
EXPECT_EQ(got.naive_blocks, 6u);
EXPECT_LT(got.stack_blocks, got.naive_blocks);
const std::uint8_t sem = static_cast<std::uint8_t>(c0 ^ c1);
const std::uint8_t want = (sem == 0) ? plain_one(heavy, {1, 1, 1})
: plain_one(light, {1, 0});
ASSERT_EQ(got.share0.size(), 1u);
EXPECT_EQ(static_cast<std::uint8_t>(got.share0[0] ^ got.share1[0]), want);
}
}
}
TEST(YaoStack, OneHotCostsTheHeaviestBranch)
{
auto a = and_n(2);
auto b = xor2();
auto c = and_n(4);
EXPECT_EQ(a.n_and(), 1u);
EXPECT_EQ(b.n_and(), 0u);
EXPECT_EQ(c.n_and(), 3u);
std::vector<netlist> branches;
branches.push_back(a);
branches.push_back(b);
branches.push_back(c);
std::vector<std::vector<std::uint8_t>> p0{
{1, 1},
{1, 0},
{1, 1, 1, 0},
};
std::vector<std::vector<std::uint8_t>> p1{
{0, 0},
{0, 1},
{0, 0, 0, 0},
};
for (std::uint16_t idx = 0; idx < 3; ++idx)
{
auto got = dpf::yao::eval_one_hot(branches, idx, 0, p0, p1);
EXPECT_EQ(got.stack_blocks, 6u);
EXPECT_EQ(got.naive_blocks, 8u);
std::vector<std::uint8_t> sem(p0[idx].size());
for (std::size_t i = 0; i < sem.size(); ++i)
sem[i] = static_cast<std::uint8_t>(p0[idx][i] ^ p1[idx][i]);
const std::uint8_t want = plain_one(branches[idx], sem);
EXPECT_EQ(static_cast<std::uint8_t>(got.share0[0] ^ got.share1[0]), want);
}
}
namespace
{
std::uint8_t xor_share(const std::vector<std::uint8_t> & a,
const std::vector<std::uint8_t> & b, std::size_t i)
{
return static_cast<std::uint8_t>(a[i] ^ b[i]);
}
netlist not1()
{
netlist nl;
nl.out(nl.not_(nl.shared_in()));
return nl;
}
netlist two_outs()
{
netlist nl;
auto a = nl.shared_in();
auto b = nl.shared_in();
nl.out(nl.and_(a, b));
nl.out(nl.xor_(a, b));
return nl;
}
netlist not_and_id()
{
netlist nl;
auto a = nl.shared_in();
nl.shared_in();
nl.out(nl.not_(a));
nl.out(a);
return nl;
}
netlist priv_and()
{
netlist nl;
auto p = nl.priv_in(0);
auto s = nl.shared_in();
nl.out(nl.and_(p, s));
return nl;
}
} // namespace
TEST(YaoStack, EveryShareSplitOfBothBranches)
{
auto then_nl = and_n(2);
auto else_nl = not1();
for (std::uint8_t c0 = 0; c0 < 2; ++c0)
{
for (std::uint8_t c1 = 0; c1 < 2; ++c1)
{
for (std::uint8_t a0 = 0; a0 < 2; ++a0)
{
for (std::uint8_t a1 = 0; a1 < 2; ++a1)
{
for (std::uint8_t b0 = 0; b0 < 2; ++b0)
{
for (std::uint8_t b1 = 0; b1 < 2; ++b1)
{
for (std::uint8_t e0 = 0; e0 < 2; ++e0)
{
for (std::uint8_t e1 = 0; e1 < 2; ++e1)
{
const std::uint8_t tp0[2] = {a0, b0};
const std::uint8_t tp1[2] = {a1, b1};
const std::uint8_t ep0[1] = {e0};
const std::uint8_t ep1[1] = {e1};
auto got = dpf::yao::eval_if(then_nl, else_nl, c0, c1,
tp0, tp1, ep0, ep1);
const std::uint8_t sem =
static_cast<std::uint8_t>(c0 ^ c1);
std::uint8_t want = 0;
if (sem == 0)
{
const std::uint8_t in[2] = {
xor_share({a0}, {a1}, 0),
static_cast<std::uint8_t>(b0 ^ b1),
};
want = plain_one(then_nl, {in[0], in[1]});
}
else
{
want = plain_one(else_nl,
{static_cast<std::uint8_t>(e0 ^ e1)});
}
EXPECT_EQ(static_cast<std::uint8_t>(
got.share0[0] ^ got.share1[0]),
want);
EXPECT_EQ(got.stack_blocks, 2u);
EXPECT_EQ(got.naive_blocks, 2u);
}
}
}
}
}
}
}
}
}
TEST(YaoStack, TwoOutputsAndAPrivateInput)
{
auto then_nl = two_outs();
auto else_nl = not_and_id();
const std::uint8_t tp0[2] = {1, 0};
const std::uint8_t tp1[2] = {1, 1};
const std::uint8_t ep0[2] = {0, 1};
const std::uint8_t ep1[2] = {1, 0};
auto got = dpf::yao::eval_if(then_nl, else_nl, 1, 0, tp0, tp1, ep0, ep1);
ASSERT_EQ(got.share0.size(), 2u);
const std::uint8_t sem_in[2] = {static_cast<std::uint8_t>(1 ^ 1),
static_cast<std::uint8_t>(0 ^ 1)};
auto want = dpf::yao::eval_plain(then_nl, sem_in);
for (std::size_t i = 0; i < 2; ++i)
EXPECT_EQ(static_cast<std::uint8_t>(got.share0[i] ^ got.share1[i]), want[i]);
auto priv = priv_and();
auto neg = not1();
const std::uint8_t pp0[2] = {1, 1};
const std::uint8_t pp1[2] = {0, 1};
const std::uint8_t np0[1] = {0};
const std::uint8_t np1[1] = {0};
auto branched = dpf::yao::eval_if(priv, neg, 0, 1, pp0, pp1, np0, np1);
EXPECT_EQ(static_cast<std::uint8_t>(branched.share0[0] ^ branched.share1[0]),
plain_one(neg, {0}));
}
TEST(YaoStack, EmptyBranchesCostNothing)
{
netlist left;
netlist right;
auto a = left.shared_in();
left.out(left.xor_(a, left.shared_in()));
auto b = right.shared_in();
right.out(right.not_(right.xor_(b, right.shared_in())));
const std::uint8_t z[2] = {1, 1};
auto got = dpf::yao::eval_if(left, right, 0, 0, z, z, z, z);
EXPECT_EQ(got.stack_blocks, 0u);
EXPECT_EQ(got.naive_blocks, 0u);
EXPECT_EQ(static_cast<std::uint8_t>(got.share0[0] ^ got.share1[0]), 0);
}
TEST(YaoStack, OneHotEveryIndexAndShare)
{
std::vector<netlist> branches;
for (unsigned i = 0; i < 8; ++i)
branches.push_back(i % 2 == 0 ? and_n(2) : xor2());
const std::uint8_t patterns[2][2] = {{0, 0}, {1, 1}};
for (std::uint16_t index = 0; index < 8; ++index)
{
for (std::uint16_t mask = 0; mask < 2; ++mask)
{
const std::uint16_t p0 = static_cast<std::uint16_t>(index ^ mask);
const std::uint16_t p1 = mask;
std::vector<std::vector<std::uint8_t>> in0(8), in1(8);
for (unsigned b = 0; b < 8; ++b)
{
in0[b] = {patterns[b & 1][0], static_cast<std::uint8_t>(b & 1u)};
in1[b] = {patterns[b & 1][1], 1};
}
auto got = dpf::yao::eval_one_hot(branches, p0, p1, in0, in1);
EXPECT_EQ(got.stack_blocks, 2u);
EXPECT_EQ(got.naive_blocks, 8u);
EXPECT_LT(got.stack_blocks, got.naive_blocks);
std::vector<std::uint8_t> sem{
static_cast<std::uint8_t>(in0[index][0] ^ in1[index][0]),
static_cast<std::uint8_t>(in0[index][1] ^ in1[index][1]),
};
EXPECT_EQ(static_cast<std::uint8_t>(got.share0[0] ^ got.share1[0]),
plain_one(branches[index], sem));
}
}
}
TEST(YaoStack, RejectsBadShapes)
{
auto a = and_n(2);
auto b = not1();
const std::uint8_t bit = 1;
EXPECT_THROW(dpf::yao::eval_if(a, b, 2, 0, &bit, &bit, &bit, &bit),
std::invalid_argument);
netlist one;
one.out(one.shared_in());
netlist two;
auto t0 = two.shared_in();
auto t1 = two.shared_in();
two.out(t0);
two.out(t1);
EXPECT_THROW(dpf::yao::eval_if(one, two, 0, 0, &bit, &bit, &bit, &bit),
std::invalid_argument);
std::vector<netlist> only{a};
EXPECT_THROW(dpf::yao::eval_one_hot(only, 0, 0, {{}}, {{}}), std::invalid_argument);
std::vector<netlist> too_many(9, a);
EXPECT_THROW(dpf::yao::eval_one_hot(too_many, 0, 0,
std::vector<std::vector<std::uint8_t>>(9),
std::vector<std::vector<std::uint8_t>>(9)),
std::invalid_argument);
std::vector<netlist> pair{a, xor2()};
EXPECT_THROW(dpf::yao::eval_one_hot(pair, 2, 0, {{1, 1}, {1, 1}}, {{0, 0}, {0, 0}}),
std::invalid_argument);
EXPECT_THROW(dpf::yao::eval_one_hot(pair, 0, 0, {{1}, {1, 1}}, {{0, 0}, {0, 0}}),
std::invalid_argument);
}