Skip to content

Commit f91bcbf

Browse files
committed
Support config reload of task durations, and reload of key and certificate.
1 parent a414153 commit f91bcbf

4 files changed

Lines changed: 99 additions & 55 deletions

File tree

cmd/server/main.go

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
package main
22

33
import (
4+
"context"
45
"flag"
56
"fmt"
67
"os"
@@ -49,6 +50,7 @@ func main() {
4950
fmt.Printf("failed to load config: %s", err)
5051
os.Exit(1)
5152
}
53+
config.CatchHUP(context.Background())
5254

5355
c, err := server.Main(config, Build, l)
5456
if err != nil {

server/command.go

Lines changed: 0 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -9,46 +9,11 @@ import (
99
"github.com/cego/nebula-provisioner/protocol"
1010
"github.com/cego/nebula-provisioner/server/store"
1111
"github.com/sirupsen/logrus"
12-
"google.golang.org/grpc"
1312
"google.golang.org/grpc/codes"
1413
"google.golang.org/grpc/status"
1514
"google.golang.org/protobuf/types/known/emptypb"
1615
)
1716

18-
func (srv *server) startUnixSocket(s *store.Store) error {
19-
srv.l.Println("Starting http unix socket")
20-
socketPath := srv.config.GetString("command.socket", "/tmp/nebula-provisioner.socket") // TODO Change default path
21-
lis, err := net.Listen("unix", socketPath)
22-
if err != nil {
23-
return err
24-
}
25-
26-
var opts []grpc.ServerOption
27-
srv.unixGrpc = grpc.NewServer(opts...)
28-
29-
c := commandServer{
30-
l: srv.l,
31-
store: s,
32-
ipManager: srv.ipManager,
33-
}
34-
protocol.RegisterServerCommandServer(srv.unixGrpc, &c)
35-
go func() {
36-
err := srv.unixGrpc.Serve(lis)
37-
if err != nil {
38-
srv.l.WithError(err).Error("Failed to start http unix socket")
39-
}
40-
}()
41-
return nil
42-
}
43-
44-
func (s *server) stopUnixSocket() error {
45-
s.l.Println("Stopping http unix socket")
46-
if s.unixGrpc != nil {
47-
s.unixGrpc.GracefulStop()
48-
}
49-
return nil
50-
}
51-
5217
type commandServer struct {
5318
protocol.UnimplementedServerCommandServer
5419

server/server.go

Lines changed: 81 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -3,10 +3,12 @@ package server
33
import (
44
"crypto/tls"
55
"fmt"
6+
"net"
67
"net/http"
78
"os"
89
"path/filepath"
910
"strings"
11+
"sync"
1012

1113
"github.com/cego/nebula-provisioner/protocol"
1214
"github.com/cego/nebula-provisioner/server/store"
@@ -17,15 +19,17 @@ import (
1719
)
1820

1921
type server struct {
20-
l *logrus.Logger
21-
config *config.C
22-
buildVersion string
23-
initialized bool
24-
store *store.Store
25-
ipManager *store.IPManager
26-
unixGrpc *grpc.Server
27-
agentService *grpc.Server
28-
tasks *tasks
22+
l *logrus.Logger
23+
config *config.C
24+
buildVersion string
25+
initialized bool
26+
store *store.Store
27+
ipManager *store.IPManager
28+
unixGrpc *grpc.Server
29+
agentService *grpc.Server
30+
tasks *tasks
31+
tlsLock sync.RWMutex
32+
tlsCertificate *tls.Certificate
2933
}
3034

3135
func Main(config *config.C, buildVersion string, logger *logrus.Logger) (*Control, error) {
@@ -34,7 +38,7 @@ func Main(config *config.C, buildVersion string, logger *logrus.Logger) (*Contro
3438
FullTimestamp: true,
3539
}
3640

37-
server := server{l, config, buildVersion, false, nil, nil, nil, nil, nil}
41+
server := server{l: l, config: config, buildVersion: buildVersion}
3842

3943
return &Control{l, server.start, server.stop, make(chan interface{})}, nil
4044
}
@@ -164,17 +168,31 @@ func (s *server) startHttpsServer(dataDir string) error {
164168

165169
tlsConfig = manager.TLSConfig()
166170
} else {
167-
168-
cert := s.config.GetString("pki.cert", "server.crt")
169-
key := s.config.GetString("pki.key", "server.key")
170-
171-
keyPair, err := tls.LoadX509KeyPair(cert, key)
171+
s.config.RegisterReloadCallback(func(_ *config.C) {
172+
s.l.Info("Reloading tls cert")
173+
keyPair, err := s.getKeyPair()
174+
if err != nil {
175+
return
176+
}
177+
s.tlsLock.Lock()
178+
defer s.tlsLock.Unlock()
179+
s.tlsCertificate = keyPair
180+
})
181+
182+
keyPair, err := s.getKeyPair()
172183
if err != nil {
173-
s.l.WithError(err).Errorf("SERVER: unable to read server key pair: %v", err)
174184
return err
175185
}
186+
s.tlsLock.Lock()
187+
defer s.tlsLock.Unlock()
188+
s.tlsCertificate = keyPair
189+
176190
tlsConfig = &tls.Config{
177-
Certificates: []tls.Certificate{keyPair},
191+
GetCertificate: func(_ *tls.ClientHelloInfo) (*tls.Certificate, error) {
192+
s.tlsLock.RLock()
193+
defer s.tlsLock.RUnlock()
194+
return s.tlsCertificate, nil
195+
},
178196
}
179197
}
180198

@@ -198,6 +216,52 @@ func (s *server) startHttpsServer(dataDir string) error {
198216
return nil
199217
}
200218

219+
func (s *server) getKeyPair() (*tls.Certificate, error) {
220+
cert := s.config.GetString("pki.cert", "server.crt")
221+
key := s.config.GetString("pki.key", "server.key")
222+
223+
keyPair, err := tls.LoadX509KeyPair(cert, key)
224+
if err != nil {
225+
s.l.WithError(err).Errorf("SERVER: unable to read server key pair: %v", err)
226+
return nil, err
227+
}
228+
return &keyPair, nil
229+
}
230+
231+
func (s *server) startUnixSocket(store *store.Store) error {
232+
s.l.Println("Starting http unix socket")
233+
socketPath := s.config.GetString("command.socket", "/tmp/nebula-provisioner.socket") // TODO Change default path
234+
lis, err := net.Listen("unix", socketPath)
235+
if err != nil {
236+
return err
237+
}
238+
239+
var opts []grpc.ServerOption
240+
s.unixGrpc = grpc.NewServer(opts...)
241+
242+
c := commandServer{
243+
l: s.l,
244+
store: store,
245+
ipManager: s.ipManager,
246+
}
247+
protocol.RegisterServerCommandServer(s.unixGrpc, &c)
248+
go func() {
249+
err := s.unixGrpc.Serve(lis)
250+
if err != nil {
251+
s.l.WithError(err).Error("Failed to start http unix socket")
252+
}
253+
}()
254+
return nil
255+
}
256+
257+
func (s *server) stopUnixSocket() error {
258+
s.l.Println("Stopping http unix socket")
259+
if s.unixGrpc != nil {
260+
s.unixGrpc.GracefulStop()
261+
}
262+
return nil
263+
}
264+
201265
func grpcHandlerFunc(g *grpc.Server, h http.Handler) http.Handler {
202266
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
203267
ct := r.Header.Get("Content-Type")

server/tasks.go

Lines changed: 16 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -23,9 +23,20 @@ func NewTasks(l *logrus.Logger, config *config.C, store *store.Store) *tasks {
2323
func (t *tasks) Start() {
2424
t.l.Infoln("Starting task scheduler")
2525

26-
renewCertTicker := time.NewTicker(t.config.GetDuration("tasks.certRenew.interval", 1*time.Hour))
27-
renewCATicker := time.NewTicker(t.config.GetDuration("tasks.caRenew.interval", 24*time.Hour))
28-
dbGCTicker := time.NewTicker(t.config.GetDuration("tasks.dbGC.interval", 5*time.Minute))
26+
renewCertDuration := func() time.Duration { return t.config.GetDuration("tasks.certRenew.interval", 1*time.Hour) }
27+
renewCADuration := func() time.Duration { return t.config.GetDuration("tasks.caRenew.interval", 24*time.Hour) }
28+
dbGCDuration := func() time.Duration { return t.config.GetDuration("tasks.dbGC.interval", 5*time.Minute) }
29+
30+
renewCertTicker := time.NewTicker(renewCertDuration())
31+
renewCATicker := time.NewTicker(renewCADuration())
32+
dbGCTicker := time.NewTicker(dbGCDuration())
33+
34+
t.config.RegisterReloadCallback(func(_ *config.C) {
35+
t.l.Info("Reloading task scheduler")
36+
renewCertTicker.Reset(renewCertDuration())
37+
renewCATicker.Reset(renewCADuration())
38+
dbGCTicker.Reset(dbGCDuration())
39+
})
2940

3041
go func() {
3142
for {
@@ -39,6 +50,8 @@ func (t *tasks) Start() {
3950
t.dbGC()
4051
case <-t.quit:
4152
renewCertTicker.Stop()
53+
renewCATicker.Stop()
54+
dbGCTicker.Stop()
4255
return
4356
}
4457
}

0 commit comments

Comments
 (0)