28c48548903dcbdf5d73438bce920af34bd66fae

Author
Harsh Mantri <24585585+cheesyhypocrisy@users.noreply.github.com>
Committer
GitHub <noreply@github.com>
Date

Message

feat: add support for certificate reloading upon SIGHUP (#710)

* feat: add support for certificate reloading upon SIGHUP

* fix: support certificate reloading for unix and add test

* fix(cmd): move cert reloader logic to the serve package

---------

Co-authored-by: Ayman Bagabas <ayman.bagabas@gmail.com>

Diff

This diff is truncated to protect this page.

  1diff --git a/cmd/soft/serve/certreloader.go b/cmd/soft/serve/certreloader.go
  2new file mode 100644
  3index 0000000000000000000000000000000000000000..34dc4d9b3bc5ab6e0ea6558ebf89a2cebb8c6b0f
  4--- /dev/null
  5+++ b/cmd/soft/serve/certreloader.go
  6@@ -0,0 +1,54 @@
  7+package serve
  8+
  9+import (
 10+	"crypto/tls"
 11+	"sync"
 12+
 13+	"charm.land/log/v2"
 14+)
 15+
 16+// CertReloader is responsible for reloading TLS certificates when a SIGHUP signal is received.
 17+type CertReloader struct {
 18+	certMu   sync.RWMutex
 19+	cert     *tls.Certificate
 20+	certPath string
 21+	keyPath  string
 22+}
 23+
 24+// NewCertReloader creates a new CertReloader that watches for SIGHUP signals.
 25+func NewCertReloader(certPath, keyPath string, logger *log.Logger) (*CertReloader, error) {
 26+	reloader := &CertReloader{
 27+		certPath: certPath,
 28+		keyPath:  keyPath,
 29+	}
 30+
 31+	cert, err := tls.LoadX509KeyPair(certPath, keyPath)
 32+	if err != nil {
 33+		return nil, err
 34+	}
 35+	reloader.cert = &cert
 36+
 37+	return reloader, nil
 38+}
 39+
 40+// Reload attempts to reload the certificate and key.
 41+func (cr *CertReloader) Reload() error {
 42+	newCert, err := tls.LoadX509KeyPair(cr.certPath, cr.keyPath)
 43+	if err != nil {
 44+		return err
 45+	}
 46+
 47+	cr.certMu.Lock()
 48+	defer cr.certMu.Unlock()
 49+	cr.cert = &newCert
 50+	return nil
 51+}
 52+
 53+// GetCertificateFunc returns a function that can be used with tls.Config.GetCertificate.
 54+func (cr *CertReloader) GetCertificateFunc() func(*tls.ClientHelloInfo) (*tls.Certificate, error) {
 55+	return func(clientHello *tls.ClientHelloInfo) (*tls.Certificate, error) {
 56+		cr.certMu.RLock()
 57+		defer cr.certMu.RUnlock()
 58+		return cr.cert, nil
 59+	}
 60+}
 61diff --git a/cmd/soft/serve/certreloader_test.go b/cmd/soft/serve/certreloader_test.go
 62new file mode 100644
 63index 0000000000000000000000000000000000000000..e22fcf790c1e80fd8c3d9827ddd57f74e1b49c06
 64--- /dev/null
 65+++ b/cmd/soft/serve/certreloader_test.go
 66@@ -0,0 +1,116 @@
 67+//go:build unix
 68+
 69+package serve
 70+
 71+import (
 72+	"crypto/rand"
 73+	"crypto/rsa"
 74+	"crypto/x509"
 75+	"crypto/x509/pkix"
 76+	"encoding/pem"
 77+	"os"
 78+	"os/signal"
 79+	"path/filepath"
 80+	"syscall"
 81+	"testing"
 82+	"time"
 83+
 84+	"charm.land/log/v2"
 85+)
 86+
 87+func generateTestCert(t *testing.T, certPath, keyPath, cn string) {
 88+	t.Helper()
 89+
 90+	privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
 91+	if err != nil {
 92+		t.Fatal(err)
 93+	}
 94+
 95+	template := x509.Certificate{
 96+		SerialNumber: nil,
 97+		Subject: pkix.Name{
 98+			CommonName: cn,
 99+		},
100+		NotBefore: time.Now(),
101+		NotAfter:  time.Now().Add(time.Hour),
102+	}
103+
104+	certBytes, err := x509.CreateCertificate(rand.Reader, &template, &template, &privateKey.PublicKey, privateKey)
105+	if err != nil {
106+		t.Fatal(err)
107+	}
108+
109+	certFile, err := os.Create(certPath)
110+	if err != nil {
111+		t.Fatal(err)
112+	}
113+	defer certFile.Close()
114+
115+	pem.Encode(certFile, &pem.Block{Type: "CERTIFICATE", Bytes: certBytes})
116+
117+	keyFile, err := os.Create(keyPath)
118+	if err != nil {
119+		t.Fatal(err)
120+	}
121+	defer keyFile.Close()
122+
123+	pem.Encode(keyFile, &pem.Block{
124+		Type:  "RSA PRIVATE KEY",
125+		Bytes: x509.MarshalPKCS1PrivateKey(privateKey),
126+	})
127+}
128+
129+func TestCertReloader(t *testing.T) {
130+	dir := t.TempDir()
131+	certPath := filepath.Join(dir, "/cert.pem")
132+	keyPath := filepath.Join(dir, "/key.pem")
133+
134+	// Initial cert
135+	generateTestCert(t, certPath, keyPath, "cert-v1")
136+
137+	logger := log.New(os.Stderr)
138+
139+	certReloader, err := NewCertReloader(certPath, keyPath, logger)
140+	if err != nil {
141+		t.Fatalf("failed to create reloader: %v", err)
142+	}
143+
144+	go func() {
145+		sigCh := make(chan os.Signal, 1)
146+		signal.Notify(sigCh, syscall.SIGHUP)
147+		for range sigCh {
148+			if err := certReloader.Reload(); err != nil {
149+				logger.Error("failed to reload certificate", "err", err)
150+			} else {
151+				logger.Info("certificate reloaded successfully")
152+			}
153+		}
154+	}()
155+
156+	getCert := certReloader.GetCertificateFunc()
157+
158+	cert1, err := getCert(nil)
159+	if err != nil {
160+		t.Fatal(err)
161+	}
162+
163+	// Replace cert on disk
164+	generateTestCert(t, certPath, keyPath, "cert-v2")
165+
166diff --git a/cmd/soft/serve/serve.go b/cmd/soft/serve/serve.go
167index 21a9f080ae16065fb5edcc32b46c0b8f2e4de441..7472f3d0be14300f5dcb827197308abfcc6d9ff6 100644
168--- a/cmd/soft/serve/serve.go
169+++ b/cmd/soft/serve/serve.go
170@@ -86,7 +86,7 @@ var (
171 			done := make(chan os.Signal, 1)
172 			doneOnce := sync.OnceFunc(func() { close(done) })
173 
174-			signal.Notify(done, os.Interrupt, syscall.SIGINT, syscall.SIGTERM)
175+			signal.Notify(done, os.Interrupt, syscall.SIGINT, syscall.SIGTERM, syscall.SIGHUP)
176 
177 			// This endpoint is added for testing purposes
178 			// It allows us to stop the server from the test suite.
179@@ -107,12 +107,23 @@ var (
180 				doneOnce()
181 			}()
182 
183-			select {
184-			case err := <-lch:
185-				if err != nil {
186-					return fmt.Errorf("server error: %w", err)
187+			for {
188+				select {
189+				case err := <-lch:
190+					if err != nil {
191+						return fmt.Errorf("server error: %w", err)
192+					}
193+				case sig := <-done:
194+					if sig == syscall.SIGHUP {
195+						s.logger.Info("received SIGHUP signal, reloading TLS certificates if enabled")
196+						if err := s.ReloadCertificates(); err != nil {
197+							s.logger.Error("failed to reload TLS certificates", "err", err)
198+						}
199+						continue
200+					}
201 				}
202-			case <-done:
203+
204+				break
205 			}
206 
207 			ctx, cancel := context.WithTimeout(ctx, 5*time.Second)
208diff --git a/cmd/soft/serve/server.go b/cmd/soft/serve/server.go
209index 4260f7576da3fccc12f5b296255f70490ed72682..fda09005097dce89960b6cd6af26431ceca54c32 100644
210--- a/cmd/soft/serve/server.go
211+++ b/cmd/soft/serve/server.go
212@@ -2,6 +2,7 @@ package serve
213 
214 import (
215 	"context"
216+	"crypto/tls"
217 	"errors"
218 	"fmt"
219 	"net/http"
220@@ -27,6 +28,7 @@ type Server struct {
221 	GitDaemon   *daemon.GitDaemon
222 	HTTPServer  *web.HTTPServer
223 	StatsServer *stats.StatsServer
224+	CertLoader  *CertReloader
225 	Cron        *cron.Scheduler
226 	Config      *config.Config
227 	Backend     *backend.Backend
228@@ -87,9 +89,28 @@ func NewServer(ctx context.Context) (*Server, error) {
229 		return nil, fmt.Errorf("create stats server: %w", err)
230 	}
231 
232+	if cfg.HTTP.TLSKeyPath != "" && cfg.HTTP.TLSCertPath != "" {
233+		srv.CertLoader, err = NewCertReloader(cfg.HTTP.TLSCertPath, cfg.HTTP.TLSKeyPath, logger)
234+		if err != nil {
235+			return nil, fmt.Errorf("create cert reloader: %w", err)
236+		}
237+
238+		srv.HTTPServer.SetTLSConfig(&tls.Config{
239+			GetCertificate: srv.CertLoader.GetCertificateFunc(),
240+		})
241+	}
242+
243 	return srv, nil
244 }
245 
246+// ReloadCertificates reloads the TLS certificates for the HTTP server.
247+func (s *Server) ReloadCertificates() error {
248+	if s.CertLoader == nil {
249+		return nil
250+	}
251+	return s.CertLoader.Reload()
252+}
253+
254 // Start starts the SSH server.
255 func (s *Server) Start() error {
256 	errg, _ := errgroup.WithContext(s.ctx)
257diff --git a/pkg/web/http.go b/pkg/web/http.go
258index 7bb255f4007d471fb83bf604357d3d60a8e192c7..531d02bfd89c4fbbf71caf64d02eaebe44585bc5 100644
259--- a/pkg/web/http.go
260+++ b/pkg/web/http.go
261@@ -2,6 +2,7 @@ package web
262 
263 import (
264 	"context"
265+	"crypto/tls"
266 	"net/http"
267 	"time"
268 
269@@ -37,6 +38,11 @@ func NewHTTPServer(ctx context.Context) (*HTTPServer, error) {
270 	return s, nil
271 }
272 
273+// SetTLSConfig sets the TLS configuration for the HTTP server.
274+func (s *HTTPServer) SetTLSConfig(tlsConfig *tls.Config) {
275+	s.Server.TLSConfig = tlsConfig
276+}
277+
278 // Close closes the HTTP server.
279 func (s *HTTPServer) Close() error {
280 	return s.Server.Close()
281@@ -44,8 +50,8 @@ func (s *HTTPServer) Close() error {
282 
283 // ListenAndServe starts the HTTP server.
284 func (s *HTTPServer) ListenAndServe() error {
285-	if s.cfg.HTTP.TLSKeyPath != "" && s.cfg.HTTP.TLSCertPath != "" {
286-		return s.Server.ListenAndServeTLS(s.cfg.HTTP.TLSCertPath, s.cfg.HTTP.TLSKeyPath)
287+	if s.Server.TLSConfig != nil {
288+		return s.Server.ListenAndServeTLS("", "")
289 	}
290 	return s.Server.ListenAndServe()
291 }