98 lines
3.3 KiB
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)
|
|
}
|
|
}
|