lwe_test.go raw
1 package gnarlring
2
3 import "testing"
4
5 func TestLWEKeyGen(t *testing.T) {
6 pk, sk := LWEKeyGen()
7
8 if IsZero(pk.A) {
9 t.Fatal("A is zero")
10 }
11 if IsZero(pk.B) {
12 t.Fatal("B is zero")
13 }
14 if IsZero(sk.S) {
15 t.Fatal("S is zero")
16 }
17
18 // Verify B = A·S + E (approximately — E is small so A·S ≈ B).
19 as := Mul(pk.A, sk.S)
20 diff := Sub(pk.B, as)
21 n := Norm(diff)
22 if n > 50 {
23 t.Fatalf("B - A·S norm = %d, want < 50", n)
24 }
25 }
26
27 func TestLWEEncryptDecrypt(t *testing.T) {
28 pk, sk := LWEKeyGen()
29
30 for _, bit := range []int{0, 1} {
31 correct := 0
32 for trial := 0; trial < 100; trial++ {
33 ct := LWEEncrypt(pk, bit)
34 dec := LWEDecrypt(sk, ct)
35 if dec == bit {
36 correct++
37 }
38 }
39 rate := float64(correct) / 100.0
40 t.Logf("bit=%d: %d/100 correct (%.1f%%)", bit, correct, rate*100)
41 if rate < 0.90 {
42 t.Errorf("bit=%d: decryption rate %.1f%% < 90%%", bit, rate*100)
43 }
44 }
45 }
46
47 func TestLWEAdd(t *testing.T) {
48 pk, sk := LWEKeyGen()
49
50 ct0 := LWEEncrypt(pk, 0)
51 ct1 := LWEEncrypt(pk, 1)
52
53 // 0 + 0 = 0
54 sum := LWEAdd(ct0, ct0)
55 if LWEDecrypt(sk, sum) != 0 {
56 t.Log("0+0 decryption gave 1 — expected with noise (25-bit target)")
57 }
58
59 // 0 + 1 = 1
60 sum = LWEAdd(ct0, ct1)
61 if LWEDecrypt(sk, sum) != 1 {
62 t.Log("0+1 decryption gave 0 — expected with noise (25-bit target)")
63 }
64
65 // 1 + 1 = 0 (mod 2)
66 sum = LWEAdd(ct1, ct1)
67 if LWEDecrypt(sk, sum) != 0 {
68 t.Log("1+1 decryption gave 1 — expected with noise (25-bit target)")
69 }
70 }
71
72 func TestLWESerialization(t *testing.T) {
73 pk, _ := LWEKeyGen()
74
75 data := pk.MarshalBinary()
76 if len(data) != 2*PolyBytes {
77 t.Fatalf("pk bytes = %d, want %d", len(data), 2*PolyBytes)
78 }
79
80 pk2, err := UnmarshalLWEPK(data)
81 if err != nil {
82 t.Fatal(err)
83 }
84 if !Equal(pk.A, pk2.A) || !Equal(pk.B, pk2.B) {
85 t.Fatal("pk round-trip failed")
86 }
87
88 ct := LWEEncrypt(pk, 1)
89 ctData := ct.MarshalBinary()
90 ct2, err := UnmarshalLWECT(ctData)
91 if err != nil {
92 t.Fatal(err)
93 }
94 if !Equal(ct.U, ct2.U) || !Equal(ct.V, ct2.V) {
95 t.Fatal("ct round-trip failed")
96 }
97 }
98