28c48548903dcbdf5d73438bce920af34bd66fae
- Author
- Harsh Mantri <24585585+cheesyhypocrisy@users.noreply.github.com>
- Committer
- GitHub <noreply@github.com>
- Date
Message
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 }