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") } }