fix: separate TLS inputs and preserve certificate reloads (#1653)

Co-authored-by: Alessandro (Ale) Segala <43508+ItalyPaleAle@users.noreply.github.com>
This commit is contained in:
Elias Schneider
2026-08-07 00:32:08 +00:00
committed by GitHub
co-authored by Alessandro Segala
parent 1f9cc5e58e
commit 06ccadfcd0
7 changed files with 300 additions and 60 deletions
+85 -43
View File
@@ -9,9 +9,9 @@ import (
"net"
"net/http"
"os"
"path/filepath"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/fsnotify/fsnotify"
@@ -264,13 +264,20 @@ func initServerProtocols() (*http.Protocols, *tls.Config, *tlsCertProvider, erro
protocols := new(http.Protocols)
protocols.SetHTTP1(true)
if common.EnvConfig.TLSCertFile == "" || common.EnvConfig.TLSKeyFile == "" {
tlsConfigured := common.EnvConfig.TLSCert != "" || common.EnvConfig.TLSKey != "" ||
common.EnvConfig.TLSCertFile != "" || common.EnvConfig.TLSKeyFile != ""
if !tlsConfigured {
protocols.SetUnencryptedHTTP2(true)
return protocols, nil, nil, nil
}
protocols.SetHTTP2(true)
certProvider, err := newCertProvider(common.EnvConfig.TLSCertFile, common.EnvConfig.TLSKeyFile)
certProvider, err := newCertProvider(
common.EnvConfig.TLSCert,
common.EnvConfig.TLSKey,
common.EnvConfig.TLSCertFile,
common.EnvConfig.TLSKeyFile,
)
if err != nil {
return nil, nil, nil, fmt.Errorf("failed to load TLS certificate: %w", err)
}
@@ -324,7 +331,7 @@ func runServer(ctx context.Context, config *serverConfig) error {
}
func startCertWatcher(ctx context.Context, certProvider *tlsCertProvider) (*fsnotify.Watcher, error) {
if certProvider == nil {
if certProvider == nil || certProvider.certFile == "" || certProvider.keyFile == "" {
return nil, nil
}
@@ -333,13 +340,18 @@ func startCertWatcher(ctx context.Context, certProvider *tlsCertProvider) (*fsno
return nil, fmt.Errorf("failed to create certificate watcher: %w", err)
}
if err := certWatcher.Add(common.EnvConfig.TLSCertFile); err != nil {
certWatcher.Close()
return nil, fmt.Errorf("failed to watch TLS certificate: %w", err)
}
if err := certWatcher.Add(common.EnvConfig.TLSKeyFile); err != nil {
certWatcher.Close()
return nil, fmt.Errorf("failed to watch TLS key: %w", err)
watchedDirectories := make(map[string]struct{}, 2)
for _, file := range []string{certProvider.certFile, certProvider.keyFile} {
directory := filepath.Dir(file)
if _, ok := watchedDirectories[directory]; ok {
continue
}
if err := certWatcher.Add(directory); err != nil {
_ = certWatcher.Close()
return nil, fmt.Errorf("failed to watch TLS directory %q: %w", directory, err)
}
watchedDirectories[directory] = struct{}{}
}
go certProvider.StartWatching(ctx, certWatcher)
@@ -348,7 +360,7 @@ func startCertWatcher(ctx context.Context, certProvider *tlsCertProvider) (*fsno
func closeCertWatcher(certWatcher *fsnotify.Watcher) {
if certWatcher != nil {
certWatcher.Close()
_ = certWatcher.Close()
}
}
@@ -362,7 +374,7 @@ func startHTTPServer(config *serverConfig) {
}
srvErr := config.server.Serve(listener)
if srvErr != http.ErrServerClosed {
if !errors.Is(srvErr, http.ErrServerClosed) {
slog.Error("Error starting app server", "error", srvErr)
os.Exit(1)
}
@@ -454,29 +466,50 @@ func enrichRequestLog(c *gin.Context, record *slog.Record) *slog.Record {
// tlsCertProvider holds certificates that can be dynamically reloaded
type tlsCertProvider struct {
certMutex sync.RWMutex
cert *tls.Certificate
certFile string
keyFile string
forceReload atomic.Bool
certMutex sync.RWMutex
cert *tls.Certificate
certFile string
keyFile string
}
// GetCertificate implements tls.GetCertificate interface for dynamic certificate loading
func (p *tlsCertProvider) GetCertificate(_ *tls.ClientHelloInfo) (*tls.Certificate, error) {
if p.forceReload.Load() {
p.certMutex.Lock()
p.forceReload.Store(false)
p.certMutex.Unlock()
}
p.certMutex.RLock()
defer p.certMutex.RUnlock()
return p.cert, nil
}
// newCertProvider creates a new certificate provider with initial certificates loaded
func newCertProvider(certFile, keyFile string) (*tlsCertProvider, error) {
cert, err := tls.LoadX509KeyPair(certFile, keyFile)
// newCertProvider creates a certificate provider from either inline data or reloadable files
func newCertProvider(certPEM, keyPEM, certFile, keyFile string) (*tlsCertProvider, error) {
inlineConfigured := certPEM != "" || keyPEM != ""
fileConfigured := certFile != "" || keyFile != ""
switch {
case inlineConfigured && fileConfigured:
return nil, errors.New("inline and file-based TLS configuration cannot be combined")
case certPEM != "" && keyPEM == "", certPEM == "" && keyPEM != "":
return nil, errors.New("inline TLS certificate and key must both be configured")
case certFile != "" && keyFile == "", certFile == "" && keyFile != "":
return nil, errors.New("TLS certificate and key files must both be configured")
case !inlineConfigured && !fileConfigured:
return nil, errors.New("TLS certificate and key must both be configured")
}
var cert tls.Certificate
var err error
if inlineConfigured {
cert, err = tls.X509KeyPair([]byte(certPEM), []byte(keyPEM))
} else {
certFile, err = filepath.Abs(certFile)
if err != nil {
return nil, fmt.Errorf("failed to resolve TLS certificate path: %w", err)
}
keyFile, err = filepath.Abs(keyFile)
if err != nil {
return nil, fmt.Errorf("failed to resolve TLS key path: %w", err)
}
cert, err = tls.LoadX509KeyPair(certFile, keyFile)
}
if err != nil {
return nil, err
}
@@ -504,9 +537,11 @@ func (p *tlsCertProvider) reloadCertificate() error {
// StartWatching begins monitoring the certificate files for changes with debouncing
func (p *tlsCertProvider) StartWatching(ctx context.Context, watcher *fsnotify.Watcher) {
debounceDuration := 1 * time.Second
const debounceDuration = time.Second
reloadTimer := time.NewTimer(debounceDuration)
reloadTimer.Stop()
defer reloadTimer.Stop()
for {
select {
@@ -516,34 +551,41 @@ func (p *tlsCertProvider) StartWatching(ctx context.Context, watcher *fsnotify.W
if !ok {
return
}
// Only process write/rename events for certificate/key files
if event.Has(fsnotify.Write | fsnotify.Rename) {
// Reset the debounce timer whenever we get a relevant event
reloadTimer.Stop()
// Drain the channel if there's a pending value
select {
case <-reloadTimer.C:
default:
}
reloadTimer.Reset(debounceDuration)
slog.Debug("TLS file change detected, debouncing", slog.String("path", event.Name))
// Ignore events that are not related to the certificate or key files
if !p.isCertificateEvent(event) {
continue
}
// Reset the debounce timer so both files can settle before the pair is reloaded
reloadTimer.Reset(debounceDuration)
slog.Debug("TLS file change detected, debouncing", slog.String("path", event.Name))
case <-reloadTimer.C:
// Timer fired - no more events in 500ms, so reload
// Reload the pair atomically after the certificate directories have settled
slog.Info("Reloading TLS certificate")
if err := p.reloadCertificate(); err != nil {
slog.Error("Failed to reload TLS certificate", "error", err)
continue
} else {
slog.Info("TLS certificate reloaded successfully")
}
p.forceReload.Store(true)
slog.Info("TLS certificate reloaded successfully")
case err, ok := <-watcher.Errors:
if !ok {
return
}
slog.Error("Certificate watcher error", "error", err)
}
}
}
func (p *tlsCertProvider) isCertificateEvent(event fsnotify.Event) bool {
if !event.Has(fsnotify.Write | fsnotify.Create | fsnotify.Rename | fsnotify.Remove) {
return false
}
eventPath := filepath.Clean(event.Name)
return eventPath == p.certFile || eventPath == p.keyFile
}
@@ -2,11 +2,22 @@ package bootstrap
import (
"bytes"
"context"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/x509"
"crypto/x509/pkix"
"encoding/json"
"encoding/pem"
"log/slog"
"math/big"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"testing"
"time"
"github.com/gin-gonic/gin"
"github.com/pocket-id/pocket-id/backend/internal/apperror"
@@ -130,3 +141,98 @@ func TestRequestLoggerLogsAtConfiguredMinimumLevel(t *testing.T) {
require.Equal(t, http.StatusNoContent, recorder.Code)
require.Contains(t, output.String(), "level=INFO")
}
func TestNewCertProviderSupportsInlineCertificateData(t *testing.T) {
certPEM, keyPEM := newTestTLSKeyPair(t, 1)
provider, err := newCertProvider(certPEM, keyPEM, "", "")
require.NoError(t, err)
require.Empty(t, provider.certFile)
require.Empty(t, provider.keyFile)
require.True(t, certProviderHasSerial(provider, 1))
watcher, err := startCertWatcher(t.Context(), provider)
require.NoError(t, err)
require.Nil(t, watcher)
}
func TestCertProviderReloadsAfterRepeatedAtomicReplacement(t *testing.T) {
tempDir := t.TempDir()
certFile := filepath.Join(tempDir, "cert.pem")
keyFile := filepath.Join(tempDir, "key.pem")
writeTestTLSKeyPair(t, certFile, keyFile, 1)
provider, err := newCertProvider("", "", certFile, keyFile)
require.NoError(t, err)
require.True(t, certProviderHasSerial(provider, 1))
ctx, cancel := context.WithCancel(t.Context())
watcher, err := startCertWatcher(ctx, provider)
require.NoError(t, err)
t.Cleanup(func() {
cancel()
closeCertWatcher(watcher)
})
for serial := int64(2); serial <= 3; serial++ {
replaceTestTLSKeyPair(t, certFile, keyFile, serial)
require.Eventually(t, func() bool {
return certProviderHasSerial(provider, serial)
}, 5*time.Second, 50*time.Millisecond)
}
}
func replaceTestTLSKeyPair(t *testing.T, certFile, keyFile string, serial int64) {
t.Helper()
replacementCertFile := certFile + ".new"
replacementKeyFile := keyFile + ".new"
writeTestTLSKeyPair(t, replacementCertFile, replacementKeyFile, serial)
require.NoError(t, os.Rename(replacementCertFile, certFile))
require.NoError(t, os.Rename(replacementKeyFile, keyFile))
}
func writeTestTLSKeyPair(t *testing.T, certFile, keyFile string, serial int64) {
t.Helper()
certPEM, keyPEM := newTestTLSKeyPair(t, serial)
require.NoError(t, os.WriteFile(certFile, []byte(certPEM), 0600))
require.NoError(t, os.WriteFile(keyFile, []byte(keyPEM), 0600))
}
func newTestTLSKeyPair(t *testing.T, serial int64) (string, string) {
t.Helper()
privateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
require.NoError(t, err)
now := time.Now()
template := &x509.Certificate{
SerialNumber: big.NewInt(serial),
Subject: pkix.Name{CommonName: "localhost"},
NotBefore: now.Add(-time.Minute),
NotAfter: now.Add(time.Hour),
KeyUsage: x509.KeyUsageDigitalSignature,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
DNSNames: []string{"localhost"},
}
certDER, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey)
require.NoError(t, err)
keyDER, err := x509.MarshalECPrivateKey(privateKey)
require.NoError(t, err)
certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER})
keyPEM := pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER})
return string(certPEM), string(keyPEM)
}
func certProviderHasSerial(provider *tlsCertProvider, serial int64) bool {
cert, err := provider.GetCertificate(nil)
if err != nil || cert == nil || len(cert.Certificate) == 0 {
return false
}
parsedCert, err := x509.ParseCertificate(cert.Certificate[0])
return err == nil && parsedCert.SerialNumber.Cmp(big.NewInt(serial)) == 0
}