proxy-pool/internal/controlplane/tlsreload/certificate_test.go

98 lines
3.3 KiB
Go

package tlsreload
import (
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"encoding/pem"
"math/big"
"os"
"path/filepath"
"testing"
"time"
)
func TestCertificateProviderReloadsPairsAndKeepsLastKnownGoodCertificate(t *testing.T) {
directory := t.TempDir()
certificatePath := filepath.Join(directory, "tls.crt")
keyPath := filepath.Join(directory, "tls.key")
writeCertificatePair(t, certificatePath, keyPath, "first")
provider, err := NewCertificateProvider(certificatePath, keyPath)
if err != nil {
t.Fatalf("NewCertificateProvider() error = %v", err)
}
certificate, err := provider.ServerCertificate(nil)
if subject := certificateSubject(t, certificate, err); subject != "first" {
t.Fatalf("initial certificate subject = %q, want first", subject)
}
writeCertificatePair(t, certificatePath, keyPath, "second")
certificate, err = provider.ClientCertificate(nil)
if subject := certificateSubject(t, certificate, err); subject != "second" {
t.Fatalf("reloaded certificate subject = %q, want second", subject)
}
if err := os.WriteFile(keyPath, []byte("incomplete"), 0o600); err != nil {
t.Fatalf("break private key: %v", err)
}
certificate, err = provider.ServerCertificate(nil)
if subject := certificateSubject(t, certificate, err); subject != "second" {
t.Fatalf("fallback certificate subject = %q, want second", subject)
}
}
func TestCertificateProviderRejectsMissingPaths(t *testing.T) {
if _, err := NewCertificateProvider("", "key"); err != ErrInvalidCertificatePaths {
t.Fatalf("NewCertificateProvider(empty certificate) error = %v, want %v", err, ErrInvalidCertificatePaths)
}
if _, err := NewCertificateProvider("certificate", ""); err != ErrInvalidCertificatePaths {
t.Fatalf("NewCertificateProvider(empty key) error = %v, want %v", err, ErrInvalidCertificatePaths)
}
}
func certificateSubject(t *testing.T, certificate *tls.Certificate, err error) string {
t.Helper()
if err != nil {
t.Fatalf("load certificate: %v", err)
}
parsed, err := x509.ParseCertificate(certificate.Certificate[0])
if err != nil {
t.Fatalf("parse certificate: %v", err)
}
return parsed.Subject.CommonName
}
func writeCertificatePair(t *testing.T, certificatePath, keyPath, commonName string) {
t.Helper()
privateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatalf("generate private key: %v", err)
}
now := time.Now()
template := x509.Certificate{
SerialNumber: big.NewInt(now.UnixNano()), Subject: pkix.Name{CommonName: commonName},
NotBefore: now.Add(-time.Minute), NotAfter: now.Add(time.Hour),
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment,
}
der, err := x509.CreateCertificate(rand.Reader, &template, &template, &privateKey.PublicKey, privateKey)
if err != nil {
t.Fatalf("create certificate: %v", err)
}
certificate := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})
privateDER, err := x509.MarshalPKCS8PrivateKey(privateKey)
if err != nil {
t.Fatalf("marshal private key: %v", err)
}
key := pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: privateDER})
if err := os.WriteFile(certificatePath, certificate, 0o600); err != nil {
t.Fatalf("write certificate: %v", err)
}
if err := os.WriteFile(keyPath, key, 0o600); err != nil {
t.Fatalf("write private key: %v", err)
}
}