141 lines
3.1 KiB
Go
141 lines
3.1 KiB
Go
package mlkem768
|
|
|
|
import (
|
|
"bytes"
|
|
"crypto/sha256"
|
|
"errors"
|
|
"testing"
|
|
)
|
|
|
|
func TestGenerateKeyPair(t *testing.T) {
|
|
adapter := New()
|
|
|
|
pub, priv, err := adapter.GenerateKeyPair(nil)
|
|
if err != nil {
|
|
t.Fatalf("GenerateKeyPair failed: %v", err)
|
|
}
|
|
|
|
if pub.SchemeID() != suiteID {
|
|
t.Errorf("pub.SchemeID() = 0x%04x, want 0x%04x", pub.SchemeID(), suiteID)
|
|
}
|
|
|
|
if priv.SchemeID() != suiteID {
|
|
t.Errorf("priv.SchemeID() = 0x%04x, want 0x%04x", priv.SchemeID(), suiteID)
|
|
}
|
|
|
|
rawPub := pub.Raw()
|
|
if len(rawPub) != 1184 {
|
|
t.Errorf("pub.Raw() len = %d, want 1184", len(rawPub))
|
|
}
|
|
|
|
rawPriv := priv.Raw()
|
|
if len(rawPriv) != 64 {
|
|
t.Errorf("priv.Raw() len = %d, want 64", len(rawPriv))
|
|
}
|
|
|
|
expectedKeyID := sha256.Sum256(rawPub[:8])
|
|
if !bytes.Equal(pub.KeyID(), expectedKeyID[:8]) {
|
|
t.Errorf("pub.KeyID() = %x, want %x", pub.KeyID(), expectedKeyID[:8])
|
|
}
|
|
|
|
if !bytes.Equal(priv.KeyID(), pub.KeyID()) {
|
|
t.Errorf("priv.KeyID() = %x, want %x", priv.KeyID(), pub.KeyID())
|
|
}
|
|
}
|
|
|
|
func TestEncapsulateReturnOrder(t *testing.T) {
|
|
adapter := New()
|
|
|
|
pub, _, err := adapter.GenerateKeyPair(nil)
|
|
if err != nil {
|
|
t.Fatalf("GenerateKeyPair failed: %v", err)
|
|
}
|
|
|
|
ct, ss, err := adapter.Encapsulate(pub, nil)
|
|
if err != nil {
|
|
t.Fatalf("Encapsulate failed: %v", err)
|
|
}
|
|
|
|
if len(ct) != 1088 {
|
|
t.Errorf("ciphertext len = %d, want 1088", len(ct))
|
|
}
|
|
|
|
if len(ss) != 32 {
|
|
t.Errorf("sharedSecret len = %d, want 32", len(ss))
|
|
}
|
|
}
|
|
|
|
func TestRoundTrip(t *testing.T) {
|
|
adapter := New()
|
|
|
|
pub, priv, err := adapter.GenerateKeyPair(nil)
|
|
if err != nil {
|
|
t.Fatalf("GenerateKeyPair failed: %v", err)
|
|
}
|
|
|
|
ct, ssEnc, err := adapter.Encapsulate(pub, nil)
|
|
if err != nil {
|
|
t.Fatalf("Encapsulate failed: %v", err)
|
|
}
|
|
|
|
ssDec, err := adapter.Decapsulate(priv, ct)
|
|
if err != nil {
|
|
t.Fatalf("Decapsulate failed: %v", err)
|
|
}
|
|
|
|
if !bytes.Equal(ssEnc, ssDec) {
|
|
t.Fatalf("shared secret mismatch: encapsulate=%x, decapsulate=%x", ssEnc, ssDec)
|
|
}
|
|
}
|
|
|
|
func TestDecapsulateTamperedCiphertext(t *testing.T) {
|
|
adapter := New()
|
|
|
|
pub, priv, err := adapter.GenerateKeyPair(nil)
|
|
if err != nil {
|
|
t.Fatalf("GenerateKeyPair failed: %v", err)
|
|
}
|
|
|
|
ct, _, err := adapter.Encapsulate(pub, nil)
|
|
if err != nil {
|
|
t.Fatalf("Encapsulate failed: %v", err)
|
|
}
|
|
|
|
ct[0] ^= 0xFF
|
|
|
|
_, err = adapter.Decapsulate(priv, ct)
|
|
if err == nil {
|
|
t.Fatal("Decapsulate with tampered ciphertext: expected error, got nil")
|
|
}
|
|
|
|
if !errors.Is(err, ErrDecapsulationFailed) {
|
|
t.Errorf("Decapsulate error = %v, want ErrDecapsulationFailed", err)
|
|
}
|
|
}
|
|
|
|
func TestRegistryRegistration(t *testing.T) {
|
|
factory, err := DefaultRegistry.Lookup(suiteID)
|
|
if err != nil {
|
|
t.Fatalf("Lookup suiteID 0x%04x failed: %v", suiteID, err)
|
|
}
|
|
|
|
instance := factory()
|
|
if instance.SchemeID() != suiteID {
|
|
t.Errorf("factory() SchemeID = 0x%04x, want 0x%04x", instance.SchemeID(), suiteID)
|
|
}
|
|
}
|
|
|
|
func TestFactoryReturnsIndependentInstances(t *testing.T) {
|
|
factory, err := DefaultRegistry.Lookup(suiteID)
|
|
if err != nil {
|
|
t.Fatalf("Lookup suiteID 0x%04x failed: %v", suiteID, err)
|
|
}
|
|
|
|
one := factory()
|
|
two := factory()
|
|
|
|
if one.SchemeID() != two.SchemeID() {
|
|
t.Error("factory() returned instances with different scheme IDs")
|
|
}
|
|
}
|