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