#include #include #include #include "dpf/arith_garble.hpp" namespace { using dpf::arith_garble::circuit; using dpf::arith_garble::wire; std::vector plain_outs(const circuit & c, const std::uint16_t * in) { auto all = dpf::arith_garble::eval_plain(c, in); std::vector out; for (auto id : c.outputs()) out.push_back(all[id]); return out; } void expect_match(const circuit & c, const std::vector & in) { auto got = dpf::arith_garble::eval_pair(c, in.data()); auto want = plain_outs(c, in.data()); ASSERT_EQ(got.opened, want); ASSERT_EQ(got.mask.size(), want.size()); for (std::size_t i = 0; i < want.size(); ++i) { EXPECT_EQ(got.opened[i], dpf::arith_garble::open_shares(got.modulus[i], got.mask[i], got.color[i])); } } } // namespace TEST(ArithGarble, ProjectionRowsAreModulusMinusOne) { circuit c; auto x = c.input(5); std::vector sq{0, 1, 4, 4, 1}; c.out(c.project(x, 5, sq)); auto got = dpf::arith_garble::eval_pair(c, std::vector{3}.data()); EXPECT_EQ(got.ciphertext_rows, 4u); EXPECT_EQ(got.opened[0], 4u); } TEST(ArithGarble, AddScaleAndPublicShift) { circuit c; auto x = c.input(7); auto y = c.input(7); auto s = c.add(x, y); auto t = c.scale(s, 3); c.out(c.add_const(t, 4)); for (std::uint16_t a = 0; a < 7; ++a) for (std::uint16_t b = 0; b < 7; ++b) expect_match(c, {a, b}); } TEST(ArithGarble, ThresholdCostsFanInRows) { circuit c; std::vector bits; for (int i = 0; i < 3; ++i) bits.push_back(c.input(4)); c.out(c.threshold(bits, 3)); auto got = dpf::arith_garble::eval_pair( c, std::vector{1, 1, 1}.data()); EXPECT_EQ(got.ciphertext_rows, 3u); EXPECT_EQ(got.opened[0], 1u); expect_match(c, {1, 1, 0}); expect_match(c, {0, 0, 0}); } TEST(ArithGarble, FanInAndOfBits) { circuit c; std::vector bits; for (int i = 0; i < 4; ++i) bits.push_back(c.input(2)); c.out(c.fanin_and(bits)); expect_match(c, {1, 1, 1, 1}); expect_match(c, {1, 1, 0, 1}); expect_match(c, {0, 0, 0, 0}); } TEST(ArithGarble, PrimeProduct) { for (std::uint16_t p : {2, 3, 5, 7}) { circuit c; auto x = c.input(p); auto y = c.input(p); c.out(c.mul(x, y)); for (std::uint16_t a = 0; a < p; ++a) for (std::uint16_t b = 0; b < p; ++b) expect_match(c, {a, b}); } } TEST(ArithGarble, ChainedProduct) { circuit c; auto x = c.input(5); auto y = c.input(5); auto z = c.input(5); c.out(c.mul(c.add(x, y), z)); expect_match(c, {2, 4, 3}); expect_match(c, {0, 1, 4}); expect_match(c, {4, 4, 0}); } TEST(ArithGarble, EveryResidueOfAProjectionAndAScale) { circuit c; auto x = c.input(11); std::vector cube(11); for (std::uint16_t i = 0; i < 11; ++i) cube[i] = static_cast((i * i * i) % 11); auto y = c.project(x, 11, cube); c.out(y); c.out(c.scale(y, 10)); c.out(c.add_const(x, 18)); for (std::uint16_t a = 0; a < 11; ++a) { expect_match(c, {a}); auto got = dpf::arith_garble::eval_pair(c, &a); EXPECT_EQ(got.ciphertext_rows, 10u); } } TEST(ArithGarble, FreeGatesSendNoRows) { circuit c; auto x = c.input(9); auto y = c.input(9); c.out(c.add_const(c.scale(c.add(x, y), 5), 4)); auto got = dpf::arith_garble::eval_pair(c, std::vector{8, 8}.data()); EXPECT_EQ(got.ciphertext_rows, 0u); EXPECT_EQ(got.opened[0], static_cast((5 * (8 + 8) + 4) % 9)); } TEST(ArithGarble, ThresholdEveryWeight) { circuit c; std::vector bits; for (int i = 0; i < 5; ++i) bits.push_back(c.input(6)); for (std::uint16_t t = 0; t <= 5; ++t) c.out(c.threshold(bits, t)); auto one = dpf::arith_garble::eval_pair( c, std::vector{1, 1, 1, 1, 1}.data()); EXPECT_EQ(one.ciphertext_rows, 5u * 6u); for (unsigned mask = 0; mask < 32; ++mask) { std::vector in(5); for (int i = 0; i < 5; ++i) in[i] = static_cast((mask >> i) & 1u); expect_match(c, in); } } TEST(ArithGarble, FanInAndEveryPattern) { circuit c; std::vector bits; for (int i = 0; i < 4; ++i) bits.push_back(c.input(2)); c.out(c.fanin_and(bits)); auto rows = dpf::arith_garble::eval_pair( c, std::vector{1, 1, 1, 1}.data()); EXPECT_EQ(rows.ciphertext_rows, 8u); for (unsigned mask = 0; mask < 16; ++mask) { std::vector in(4); for (int i = 0; i < 4; ++i) in[i] = static_cast((mask >> i) & 1u); expect_match(c, in); } } TEST(ArithGarble, PrimeElevenAndBitScale) { circuit mul; auto x = mul.input(11); auto y = mul.input(11); mul.out(mul.mul(x, y)); for (std::uint16_t a = 0; a < 11; ++a) for (std::uint16_t b = 0; b < 11; ++b) expect_match(mul, {a, b}); circuit pass; auto word = pass.input(5); auto bit = pass.input(2); pass.out(pass.bit_scale(word, bit)); for (std::uint16_t w = 0; w < 5; ++w) for (std::uint16_t b = 0; b < 2; ++b) expect_match(pass, {w, b}); } TEST(ArithGarble, TwoEvalutionsOpenTheSameValue) { circuit c; auto x = c.input(5); auto y = c.input(5); c.out(c.mul(x, y)); const std::uint16_t in[2] = {3, 4}; auto a = dpf::arith_garble::eval_pair(c, in); auto b = dpf::arith_garble::eval_pair(c, in); EXPECT_EQ(a.opened, b.opened); EXPECT_EQ(a.opened[0], static_cast((3 * 4) % 5)); } TEST(ArithGarble, RejectsBadModuliTablesAndInputs) { circuit c; EXPECT_THROW(c.input(1), std::invalid_argument); EXPECT_THROW(c.input(129), std::invalid_argument); auto x = c.input(8); auto y = c.input(7); EXPECT_THROW(c.add(x, y), std::invalid_argument); EXPECT_THROW(c.scale(x, 2), std::invalid_argument); EXPECT_THROW(c.scale(x, 0), std::invalid_argument); EXPECT_THROW(c.project(x, 3, {0, 1}), std::invalid_argument); EXPECT_THROW(c.project(x, 3, std::vector(8, 3)), std::invalid_argument); auto bit = c.input(2); EXPECT_THROW(c.bit_scale(x, x), std::invalid_argument); EXPECT_THROW(c.fanin_and({x}), std::invalid_argument); EXPECT_THROW(c.threshold({}, 0), std::invalid_argument); circuit bare; bare.input(3); EXPECT_THROW(dpf::arith_garble::eval_pair(bare, std::vector{0}.data()), std::invalid_argument); circuit one; one.out(one.input(4)); EXPECT_THROW(expect_match(one, {4}), std::invalid_argument); (void)bit; }