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