Files

363 lines
8.0 KiB
Go

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