package keymanager import ( "bytes" "crypto/rand" "crypto/sha256" "encoding/pem" "errors" "os" "path/filepath" "testing" "git.tswf.io/infra/go-synapse-backupper/pkg/adapters/crypto/mlkem768" "git.tswf.io/infra/go-synapse-backupper/pkg/adapters/crypto/x25519" "git.tswf.io/infra/go-synapse-backupper/pkg/domain/crypto" ) func makeRegistry( t *testing.T, ) crypto.Registry { reg := crypto.NewRegistry() if err := reg.Register( 0x0006, func() crypto.KEM { return mlkem768.New() }, ); err != nil { t.Fatalf("register mlkem768: %v", err) } if err := reg.Register( 0x0007, func() crypto.KEM { return x25519.New() }, ); err != nil { t.Fatalf("register x25519: %v", err) } return reg } func TestGenerateMLKEM768(t *testing.T) { reg := makeRegistry(t) km := NewKeyManager(reg) var pubOut, privOut bytes.Buffer err := km.Generate( 0x0006, &pubOut, &privOut, rand.Reader, ) if err != nil { t.Fatalf("Generate failed: %v", err) } pubBlock, _ := pem.Decode(pubOut.Bytes()) if pubBlock == nil { t.Fatal("failed to decode public key PEM") } if pubBlock.Type != "ML-KEM-768 PUBLIC KEY" { t.Errorf("pub PEM type = %q, want %q", pubBlock.Type, "ML-KEM-768 PUBLIC KEY") } if len(pubBlock.Bytes) != 1184 { t.Errorf("pub raw len = %d, want 1184", len(pubBlock.Bytes)) } privBlock, _ := pem.Decode(privOut.Bytes()) if privBlock == nil { t.Fatal("failed to decode private key PEM") } if privBlock.Type != "ML-KEM-768 PRIVATE KEY" { t.Errorf("priv PEM type = %q, want %q", privBlock.Type, "ML-KEM-768 PRIVATE KEY") } if len(privBlock.Bytes) != 64 { t.Errorf("priv raw len = %d, want 64", len(privBlock.Bytes)) } } func TestGenerateX25519(t *testing.T) { reg := makeRegistry(t) km := NewKeyManager(reg) var pubOut, privOut bytes.Buffer err := km.Generate( 0x0007, &pubOut, &privOut, rand.Reader, ) if err != nil { t.Fatalf("Generate failed: %v", err) } pubBlock, _ := pem.Decode(pubOut.Bytes()) if pubBlock == nil { t.Fatal("failed to decode public key PEM") } if pubBlock.Type != "X25519 PUBLIC KEY" { t.Errorf("pub PEM type = %q, want %q", pubBlock.Type, "X25519 PUBLIC KEY") } if len(pubBlock.Bytes) != 32 { t.Errorf("pub raw len = %d, want 32", len(pubBlock.Bytes)) } privBlock, _ := pem.Decode(privOut.Bytes()) if privBlock == nil { t.Fatal("failed to decode private key PEM") } if privBlock.Type != "X25519 PRIVATE KEY" { t.Errorf("priv PEM type = %q, want %q", privBlock.Type, "X25519 PRIVATE KEY") } if len(privBlock.Bytes) != 32 { t.Errorf("priv raw len = %d, want 32", len(privBlock.Bytes)) } } func TestLoadPubMLKEM768(t *testing.T) { reg := makeRegistry(t) km := NewKeyManager(reg) var pubOut, privOut bytes.Buffer err := km.Generate( 0x0006, &pubOut, &privOut, rand.Reader, ) if err != nil { t.Fatalf("Generate failed: %v", err) } dir := t.TempDir() pubPath := filepath.Join(dir, "test.pub.pem") if err := os.WriteFile(pubPath, pubOut.Bytes(), 0o644); err != nil { t.Fatalf("write pub file: %v", err) } pub, err := km.LoadPub(pubPath, 0x0006) if err != nil { t.Fatalf("LoadPub failed: %v", err) } if pub.SchemeID() != 0x0006 { t.Errorf("pub.SchemeID() = 0x%04x, want 0x0006", pub.SchemeID()) } if len(pub.Raw()) != 1184 { t.Errorf("pub.Raw() len = %d, want 1184", len(pub.Raw())) } expectedKeyID := sha256.Sum256(pub.Raw()[:8]) if !bytes.Equal(pub.KeyID(), expectedKeyID[:8]) { t.Errorf("pub.KeyID() = %x, want %x", pub.KeyID(), expectedKeyID[:8]) } } func TestLoadPrivMLKEM768(t *testing.T) { reg := makeRegistry(t) km := NewKeyManager(reg) var pubOut, privOut bytes.Buffer err := km.Generate( 0x0006, &pubOut, &privOut, rand.Reader, ) if err != nil { t.Fatalf("Generate failed: %v", err) } dir := t.TempDir() privPath := filepath.Join(dir, "test.priv.pem") if err := os.WriteFile(privPath, privOut.Bytes(), 0o600); err != nil { t.Fatalf("write priv file: %v", err) } priv, err := km.LoadPriv(privPath, 0x0006) if err != nil { t.Fatalf("LoadPriv failed: %v", err) } if priv.SchemeID() != 0x0006 { t.Errorf("priv.SchemeID() = 0x%04x, want 0x0006", priv.SchemeID()) } if len(priv.Raw()) != 64 { t.Errorf("priv.Raw() len = %d, want 64", len(priv.Raw())) } } func TestLoadPubX25519(t *testing.T) { reg := makeRegistry(t) km := NewKeyManager(reg) var pubOut, privOut bytes.Buffer err := km.Generate( 0x0007, &pubOut, &privOut, rand.Reader, ) if err != nil { t.Fatalf("Generate failed: %v", err) } dir := t.TempDir() pubPath := filepath.Join(dir, "test.pub.pem") if err := os.WriteFile(pubPath, pubOut.Bytes(), 0o644); err != nil { t.Fatalf("write pub file: %v", err) } pub, err := km.LoadPub(pubPath, 0x0007) if err != nil { t.Fatalf("LoadPub failed: %v", err) } if pub.SchemeID() != 0x0007 { t.Errorf("pub.SchemeID() = 0x%04x, want 0x0007", pub.SchemeID()) } if len(pub.Raw()) != 32 { t.Errorf("pub.Raw() len = %d, want 32", len(pub.Raw())) } expectedKeyID := sha256.Sum256(pub.Raw()) if !bytes.Equal(pub.KeyID(), expectedKeyID[:8]) { t.Errorf("pub.KeyID() = %x, want %x", pub.KeyID(), expectedKeyID[:8]) } } func TestLoadPrivX25519(t *testing.T) { reg := makeRegistry(t) km := NewKeyManager(reg) var pubOut, privOut bytes.Buffer err := km.Generate( 0x0007, &pubOut, &privOut, rand.Reader, ) if err != nil { t.Fatalf("Generate failed: %v", err) } dir := t.TempDir() privPath := filepath.Join(dir, "test.priv.pem") if err := os.WriteFile(privPath, privOut.Bytes(), 0o600); err != nil { t.Fatalf("write priv file: %v", err) } priv, err := km.LoadPriv(privPath, 0x0007) if err != nil { t.Fatalf("LoadPriv failed: %v", err) } if priv.SchemeID() != 0x0007 { t.Errorf("priv.SchemeID() = 0x%04x, want 0x0007", priv.SchemeID()) } if len(priv.Raw()) != 32 { t.Errorf("priv.Raw() len = %d, want 32", len(priv.Raw())) } } func TestLoadPubWrongPEMType(t *testing.T) { reg := makeRegistry(t) km := NewKeyManager(reg) dir := t.TempDir() pubPath := filepath.Join(dir, "wrong.pub.pem") block := &pem.Block{ Type: "X25519 PUBLIC KEY", Bytes: make([]byte, 32), } data := pem.EncodeToMemory(block) if err := os.WriteFile(pubPath, data, 0o644); err != nil { t.Fatalf("write pub file: %v", err) } _, err := km.LoadPub(pubPath, 0x0006) if err == nil { t.Fatal("expected error for wrong PEM type") } if !errors.Is(err, ErrPEMTypeMismatch) { t.Errorf("error = %v, want ErrPEMTypeMismatch", err) } } func TestLoadPrivWrongPEMType(t *testing.T) { reg := makeRegistry(t) km := NewKeyManager(reg) dir := t.TempDir() privPath := filepath.Join(dir, "wrong.priv.pem") block := &pem.Block{ Type: "ML-KEM-768 PRIVATE KEY", Bytes: make([]byte, 64), } data := pem.EncodeToMemory(block) if err := os.WriteFile(privPath, data, 0o600); err != nil { t.Fatalf("write priv file: %v", err) } _, err := km.LoadPriv(privPath, 0x0007) if err == nil { t.Fatal("expected error for wrong PEM type") } if !errors.Is(err, ErrPEMTypeMismatch) { t.Errorf("error = %v, want ErrPEMTypeMismatch", err) } } func TestLoadPubTruncatedPEM(t *testing.T) { reg := makeRegistry(t) km := NewKeyManager(reg) dir := t.TempDir() pubPath := filepath.Join(dir, "truncated.pub.pem") if err := os.WriteFile(pubPath, []byte("-----BEGIN ML-KEM-768 PUBLIC KEY-----\nnotbase64\n"), 0o644); err != nil { t.Fatalf("write pub file: %v", err) } _, err := km.LoadPub(pubPath, 0x0006) if err == nil { t.Fatal("expected error for truncated PEM") } if !errors.Is(err, ErrInvalidPEM) { t.Errorf("error = %v, want ErrInvalidPEM", err) } } func TestLoadPrivTruncatedPEM(t *testing.T) { reg := makeRegistry(t) km := NewKeyManager(reg) dir := t.TempDir() privPath := filepath.Join(dir, "truncated.priv.pem") if err := os.WriteFile(privPath, []byte("-----BEGIN X25519 PRIVATE KEY-----\nnotbase64\n"), 0o600); err != nil { t.Fatalf("write priv file: %v", err) } _, err := km.LoadPriv(privPath, 0x0007) if err == nil { t.Fatal("expected error for truncated PEM") } if !errors.Is(err, ErrInvalidPEM) { t.Errorf("error = %v, want ErrInvalidPEM", err) } }