Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion docs/Webhook-Parameters.md
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@ Usage of webhook:
-hooks value
path to the json file containing defined hooks the webhook should serve, use multiple times to load from different files
-hotreload
watch hooks file for changes and reload them automatically
watch hooks file for changes and reload them automatically; with -secure, also reload the certificate and key files when they change
-http-methods string
set default allowed HTTP methods (ie. "POST"); separate methods with comma
-ip string
Expand Down
77 changes: 77 additions & 0 deletions tls.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,10 @@ import (
"crypto/tls"
"io"
"log"
"os"
"strings"
"sync"
"time"
)

func writeTLSSupportedCipherStrings(w io.Writer, min uint16) error {
Expand Down Expand Up @@ -83,3 +86,77 @@ func getTLSCipherSuites(v string) []uint16 {

return suites
}

// certReloader serves the TLS certificate from disk and reloads it whenever
// the certificate or key file changes.
type certReloader struct {
certPath string
keyPath string

mu sync.Mutex
cert *tls.Certificate
certModTime time.Time
keyModTime time.Time
}

func newCertReloader(certPath, keyPath string) (*certReloader, error) {
r := &certReloader{certPath: certPath, keyPath: keyPath}

if err := r.load(); err != nil {
return nil, err
}

return r, nil
}

func (r *certReloader) load() error {
certInfo, err := os.Stat(r.certPath)
if err != nil {
return err
}

keyInfo, err := os.Stat(r.keyPath)
if err != nil {
return err
}

cert, err := tls.LoadX509KeyPair(r.certPath, r.keyPath)
if err != nil {
return err
}

r.cert = &cert
r.certModTime = certInfo.ModTime()
r.keyModTime = keyInfo.ModTime()

return nil
}

func (r *certReloader) changed() bool {
certInfo, err := os.Stat(r.certPath)
if err != nil {
return false
}

keyInfo, err := os.Stat(r.keyPath)
if err != nil {
return false
}

return !certInfo.ModTime().Equal(r.certModTime) || !keyInfo.ModTime().Equal(r.keyModTime)
}

func (r *certReloader) getCertificate(*tls.ClientHelloInfo) (*tls.Certificate, error) {
r.mu.Lock()
defer r.mu.Unlock()

if r.changed() {
if err := r.load(); err != nil {
log.Printf("error reloading certificate %s: %v\n", r.certPath, err)
} else {
log.Printf("certificate %s reloaded\n", r.certPath)
}
}

return r.cert, nil
}
118 changes: 118 additions & 0 deletions tls_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,118 @@
package main

import (
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"encoding/pem"
"math/big"
"net"
"os"
"path/filepath"
"testing"
"time"
)

func writeSelfSignedCert(t *testing.T, certPath, keyPath string, serial int64) {
t.Helper()

priv, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatal(err)
}

tmpl := &x509.Certificate{
SerialNumber: big.NewInt(serial),
Subject: pkix.Name{CommonName: "webhook-test"},
NotBefore: time.Now().Add(-time.Hour),
NotAfter: time.Now().Add(time.Hour),
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
IPAddresses: []net.IP{net.IPv4(127, 0, 0, 1), net.IPv6loopback},
}

der, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, &priv.PublicKey, priv)
if err != nil {
t.Fatal(err)
}

keyDER, err := x509.MarshalECPrivateKey(priv)
if err != nil {
t.Fatal(err)
}

certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})
keyPEM := pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER})

if err := os.WriteFile(certPath, certPEM, 0600); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(keyPath, keyPEM, 0600); err != nil {
t.Fatal(err)
}

mtime := time.Unix(1_000_000_000+serial, 0)
if err := os.Chtimes(certPath, mtime, mtime); err != nil {
t.Fatal(err)
}
if err := os.Chtimes(keyPath, mtime, mtime); err != nil {
t.Fatal(err)
}
}

func certSerial(t *testing.T, cert *tls.Certificate) int64 {
t.Helper()

leaf, err := x509.ParseCertificate(cert.Certificate[0])
if err != nil {
t.Fatal(err)
}

return leaf.SerialNumber.Int64()
}

func TestCertReloader(t *testing.T) {
dir := t.TempDir()
certPath := filepath.Join(dir, "cert.pem")
keyPath := filepath.Join(dir, "key.pem")

writeSelfSignedCert(t, certPath, keyPath, 1)

reloader, err := newCertReloader(certPath, keyPath)
if err != nil {
t.Fatal(err)
}

cert, err := reloader.getCertificate(nil)
if err != nil {
t.Fatal(err)
}
if got := certSerial(t, cert); got != 1 {
t.Fatalf("serial = %d, want 1", got)
}

writeSelfSignedCert(t, certPath, keyPath, 2)

cert, err = reloader.getCertificate(nil)
if err != nil {
t.Fatal(err)
}
if got := certSerial(t, cert); got != 2 {
t.Fatalf("serial after renewal = %d, want 2", got)
}

if err := os.Remove(keyPath); err != nil {
t.Fatal(err)
}

cert, err = reloader.getCertificate(nil)
if err != nil {
t.Fatal(err)
}
if got := certSerial(t, cert); got != 2 {
t.Fatalf("serial with missing key = %d, want previous certificate 2", got)
}
}
16 changes: 14 additions & 2 deletions webhook.go
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,7 @@ var (
logPath = flag.String("logfile", "", "send log output to a file; implicitly enables verbose logging")
debug = flag.Bool("debug", false, "show debug output")
noPanic = flag.Bool("nopanic", false, "do not panic if hooks cannot be loaded when webhook is not running in verbose mode")
hotReload = flag.Bool("hotreload", false, "watch hooks file for changes and reload them automatically")
hotReload = flag.Bool("hotreload", false, "watch hooks file for changes and reload them automatically; with -secure, also reload the certificate and key files when they change")
hooksURLPrefix = flag.String("urlprefix", "hooks", "url prefix to use for served hooks (protocol://yourserver:port/PREFIX/:hook-id)")
secure = flag.Bool("secure", false, "use HTTPS instead of HTTP")
asTemplate = flag.Bool("template", false, "parse hooks file as a Go template")
Expand Down Expand Up @@ -310,7 +310,19 @@ func main() {
svr.TLSNextProto = make(map[string]func(*http.Server, *tls.Conn, http.Handler)) // disable http/2

log.Printf("serving hooks on https://%s%s", addr, makeHumanPattern(hooksURLPrefix))
log.Print(svr.ServeTLS(ln, *cert, *key))

if !*hotReload {
log.Print(svr.ServeTLS(ln, *cert, *key))
return
}

reloader, err := newCertReloader(*cert, *key)
if err != nil {
log.Fatal("error loading certificate\n", err)
}

svr.TLSConfig.GetCertificate = reloader.getCertificate
log.Print(svr.ServeTLS(ln, "", ""))
}

func hookHandler(w http.ResponseWriter, r *http.Request) {
Expand Down
Loading