#include #include #include #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 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 & 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 h0{1, 1, 1}; std::vector h1{0, 0, 0}; std::vector l0{1, 1}; std::vector 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(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(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 branches; branches.push_back(a); branches.push_back(b); branches.push_back(c); std::vector> p0{ {1, 1}, {1, 0}, {1, 1, 1, 0}, }; std::vector> 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 sem(p0[idx].size()); for (std::size_t i = 0; i < sem.size(); ++i) sem[i] = static_cast(p0[idx][i] ^ p1[idx][i]); const std::uint8_t want = plain_one(branches[idx], sem); EXPECT_EQ(static_cast(got.share0[0] ^ got.share1[0]), want); } } namespace { std::uint8_t xor_share(const std::vector & a, const std::vector & b, std::size_t i) { return static_cast(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(c0 ^ c1); std::uint8_t want = 0; if (sem == 0) { const std::uint8_t in[2] = { xor_share({a0}, {a1}, 0), static_cast(b0 ^ b1), }; want = plain_one(then_nl, {in[0], in[1]}); } else { want = plain_one(else_nl, {static_cast(e0 ^ e1)}); } EXPECT_EQ(static_cast( 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(1 ^ 1), static_cast(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(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(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(got.share0[0] ^ got.share1[0]), 0); } TEST(YaoStack, OneHotEveryIndexAndShare) { std::vector 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(index ^ mask); const std::uint16_t p1 = mask; std::vector> in0(8), in1(8); for (unsigned b = 0; b < 8; ++b) { in0[b] = {patterns[b & 1][0], static_cast(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 sem{ static_cast(in0[index][0] ^ in1[index][0]), static_cast(in0[index][1] ^ in1[index][1]), }; EXPECT_EQ(static_cast(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 only{a}; EXPECT_THROW(dpf::yao::eval_one_hot(only, 0, 0, {{}}, {{}}), std::invalid_argument); std::vector too_many(9, a); EXPECT_THROW(dpf::yao::eval_one_hot(too_many, 0, 0, std::vector>(9), std::vector>(9)), std::invalid_argument); std::vector 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); }