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>
241 lines
6.8 KiB
C++
241 lines
6.8 KiB
C++
#include <gtest/gtest.h>
|
|
|
|
#include <cstdint>
|
|
#include <vector>
|
|
|
|
#include "dpf/arith_garble.hpp"
|
|
|
|
namespace
|
|
{
|
|
|
|
using dpf::arith_garble::circuit;
|
|
using dpf::arith_garble::wire;
|
|
|
|
std::vector<std::uint16_t> plain_outs(const circuit & c,
|
|
const std::uint16_t * in)
|
|
{
|
|
auto all = dpf::arith_garble::eval_plain(c, in);
|
|
std::vector<std::uint16_t> out;
|
|
for (auto id : c.outputs())
|
|
out.push_back(all[id]);
|
|
return out;
|
|
}
|
|
|
|
void expect_match(const circuit & c, const std::vector<std::uint16_t> & 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<std::uint16_t> sq{0, 1, 4, 4, 1};
|
|
c.out(c.project(x, 5, sq));
|
|
auto got = dpf::arith_garble::eval_pair(c, std::vector<std::uint16_t>{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<wire> 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<std::uint16_t>{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<wire> 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<std::uint16_t> cube(11);
|
|
for (std::uint16_t i = 0; i < 11; ++i)
|
|
cube[i] = static_cast<std::uint16_t>((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<std::uint16_t>{8, 8}.data());
|
|
EXPECT_EQ(got.ciphertext_rows, 0u);
|
|
EXPECT_EQ(got.opened[0], static_cast<std::uint16_t>((5 * (8 + 8) + 4) % 9));
|
|
}
|
|
|
|
TEST(ArithGarble, ThresholdEveryWeight)
|
|
{
|
|
circuit c;
|
|
std::vector<wire> 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<std::uint16_t>{1, 1, 1, 1, 1}.data());
|
|
EXPECT_EQ(one.ciphertext_rows, 5u * 6u);
|
|
for (unsigned mask = 0; mask < 32; ++mask)
|
|
{
|
|
std::vector<std::uint16_t> in(5);
|
|
for (int i = 0; i < 5; ++i)
|
|
in[i] = static_cast<std::uint16_t>((mask >> i) & 1u);
|
|
expect_match(c, in);
|
|
}
|
|
}
|
|
|
|
TEST(ArithGarble, FanInAndEveryPattern)
|
|
{
|
|
circuit c;
|
|
std::vector<wire> 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<std::uint16_t>{1, 1, 1, 1}.data());
|
|
EXPECT_EQ(rows.ciphertext_rows, 8u);
|
|
for (unsigned mask = 0; mask < 16; ++mask)
|
|
{
|
|
std::vector<std::uint16_t> in(4);
|
|
for (int i = 0; i < 4; ++i)
|
|
in[i] = static_cast<std::uint16_t>((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<std::uint16_t>((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<std::uint16_t>(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<std::uint16_t>{0}.data()),
|
|
std::invalid_argument);
|
|
circuit one;
|
|
one.out(one.input(4));
|
|
EXPECT_THROW(expect_match(one, {4}), std::invalid_argument);
|
|
(void)bit;
|
|
}
|