added raft over tls

This commit is contained in:
ston1th 2020-11-15 00:30:35 +01:00
commit 585baea183
16 changed files with 563 additions and 300 deletions

36
pkg/raft/cluster.go Normal file
View file

@ -0,0 +1,36 @@
package raft
import (
"errors"
"git.giftfish.de/ston1th/vipman/pkg/cluster"
"github.com/hashicorp/raft"
)
type Cluster struct {
kv *KV
r *raft.Raft
stepdown chan struct{}
stop chan struct{}
done chan struct{}
Callbacks cluster.Callbacks
}
func NewCluster(callbacks cluster.Callbacks) (cluster.Cluster, error) {
if callbacks.Leader == nil {
return nil, errors.New("Leader is nil")
}
if callbacks.Follower == nil {
return nil, errors.New("Follower is nil")
}
return &Cluster{
kv: &KV{
m: make(map[string][]byte),
},
stepdown: make(chan struct{}, 1),
stop: make(chan struct{}, 1),
done: make(chan struct{}, 1),
Callbacks: callbacks,
}, nil
}

99
pkg/raft/fsm.go Normal file
View file

@ -0,0 +1,99 @@
package raft
import (
"fmt"
"io"
"time"
"git.giftfish.de/ston1th/raftbbolt/msgpack"
"github.com/hashicorp/raft"
)
const (
retainSnapshotCount = 2
raftTimeout = 10 * time.Second
)
type Op uint8
const (
SET Op = iota
DELETE
)
type raftCmd struct {
Op Op
K string
V []byte
}
type fsm KV
func (f *fsm) Apply(l *raft.Log) interface{} {
var c raftCmd
if err := msgpack.Unmarshal(l.Data, &c); err != nil {
panic(fmt.Sprintf("failed to unmarshal command: %s", err.Error()))
}
switch c.Op {
case SET:
f.applySet(c.K, c.V)
case DELETE:
f.applyDelete(c.K)
default:
panic(fmt.Sprintf("unrecognized command op: %s", c.Op))
}
return nil
}
func (f *fsm) Snapshot() (raft.FSMSnapshot, error) {
f.mu.Lock()
defer f.mu.Unlock()
m := make(kv)
for k, v := range f.m {
m[k] = v
}
return &fsmSnapshot{m: m}, nil
}
func (f *fsm) Restore(rc io.ReadCloser) error {
o := make(kv)
if err := msgpack.NewDecoder(rc).Decode(&o); err != nil {
return err
}
f.m = o
return nil
}
func (f *fsm) applySet(k string, v []byte) {
f.mu.Lock()
f.m[k] = v
f.mu.Unlock()
}
func (f *fsm) applyDelete(k string) {
f.mu.Lock()
delete(f.m, k)
f.mu.Unlock()
}
type fsmSnapshot struct {
m kv
}
func (f *fsmSnapshot) Persist(sink raft.SnapshotSink) error {
err := func() error {
err := msgpack.NewEncoder(sink).Encode(f.m)
if err != nil {
return err
}
return sink.Close()
}()
if err != nil {
sink.Cancel()
}
return err
}
func (f *fsmSnapshot) Release() {}

56
pkg/raft/kv.go Normal file
View file

@ -0,0 +1,56 @@
package raft
import (
"errors"
"sync"
"git.giftfish.de/ston1th/raftbbolt/msgpack"
"github.com/hashicorp/raft"
)
var ErrNotLeader = errors.New("not leader")
type kv map[string][]byte
type KV struct {
mu sync.Mutex
m kv
r *raft.Raft
}
func (kv *KV) Get(k string) []byte {
kv.mu.Lock()
defer kv.mu.Unlock()
return kv.m[k]
}
func (kv *KV) Set(k string, v []byte) error {
if kv.r.State() != raft.Leader {
return ErrNotLeader
}
c := &raftCmd{SET, k, v}
b, err := msgpack.Marshal(c)
if err != nil {
return err
}
f := kv.r.Apply(b, raftTimeout)
return f.Error()
}
// Delete deletes the given key.
func (kv *KV) Delete(k string) error {
if kv.r.State() != raft.Leader {
return ErrNotLeader
}
c := &raftCmd{Op: DELETE, K: k}
b, err := msgpack.Marshal(c)
if err != nil {
return err
}
f := kv.r.Apply(b, raftTimeout)
return f.Error()
}

235
pkg/raft/raft.go Normal file
View file

@ -0,0 +1,235 @@
package raft
import (
"context"
"fmt"
"net"
"net/http"
"os"
"path/filepath"
"runtime"
"time"
"git.giftfish.de/ston1th/raftbbolt"
"git.giftfish.de/ston1th/vipman/pkg/config"
"git.giftfish.de/ston1th/vipman/pkg/raft/tls"
"github.com/go-logr/logr"
hclog "github.com/hashicorp/go-hclog"
"github.com/hashicorp/raft"
"k8s.io/klog/v2"
)
type raftLog string
func (r *raftLog) Write(p []byte) (n int, err error) {
return os.Stdout.Write(append([]byte(*r), p...))
}
func (c *Cluster) Start(log logr.Logger, raftcfg *config.Config) error {
cfg := raftcfg.Cluster.Raft
hcl := hclog.New(&hclog.LoggerOptions{
Name: "raft",
Level: hclog.LevelFromString(cfg.Level),
Output: hclog.DefaultOutput,
TimeFormat: "I0102 15:04:05.000000",
})
rc := raft.DefaultConfig()
rc.ProtocolVersion = 3
rc.LocalID = raft.ServerID(cfg.ID)
rc.Logger = hcl
pool := 3
timeout := 10 * time.Second
var transport *raft.NetworkTransport
if cfg.TLS != nil {
log.Info("transport: tls")
s, err := tls.NewStream(cfg)
if err != nil {
return err
}
transport = raft.NewNetworkTransportWithLogger(s, pool, timeout, hcl)
} else {
log.Info("transport: tcp")
address, err := net.ResolveTCPAddr("tcp", cfg.Address)
if err != nil {
return err
}
transport, err = raft.NewTCPTransportWithLogger(cfg.Address, address, pool, timeout, hcl)
if err != nil {
return err
}
}
// Create Raft structures
//snapshots := raft.NewInmemSnapshotStore()
snapshots, err := raft.NewFileSnapshotStoreWithLogger(cfg.Dir, retainSnapshotCount, hcl)
if err != nil {
return fmt.Errorf("file snapshot store: %s", err)
}
store := filepath.Join(cfg.Dir, "store.db")
//logStore := raft.NewInmemStore()
bootstrap := true
if _, err := os.Stat(store); err == nil {
bootstrap = false
}
stableStore, err := raftbbolt.NewStore(store)
if err != nil {
return err
}
logStore := stableStore
// Cluster configuration
configuration := raft.Configuration{}
// Add Local Peer
configuration.Servers = append(configuration.Servers, raft.Server{
ID: raft.ServerID(cfg.ID),
Address: raft.ServerAddress(cfg.Address),
})
for _, p := range cfg.Peers {
if cfg.Address != p.Address {
configuration.Servers = append(configuration.Servers, raft.Server{
ID: raft.ServerID(p.ID),
Address: raft.ServerAddress(p.Address)})
}
}
if bootstrap {
if err := raft.BootstrapCluster(rc, logStore, stableStore, snapshots, transport, configuration); err != nil {
return err
}
}
// Create RAFT instance
c.r, err = raft.NewRaft(rc, (*fsm)(c.kv), logStore, stableStore, snapshots, transport)
if err != nil {
return err
}
c.kv.r = c.r
t := time.NewTicker(time.Second)
time.Sleep(time.Second * 3)
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
leading := false
following := false
for {
if !leading && cfg.Address == string(c.r.Leader()) {
leading = true
go c.leader(ctx)
} else if !following && !leading {
following = true
go c.follower(ctx)
}
select {
case leader := <-c.r.LeaderCh():
if leader {
if following {
cancel()
following = false
ctx, cancel = context.WithCancel(context.Background())
}
if !leading {
leading = true
go c.leader(ctx)
}
} else if c.r.State() == raft.Follower {
if leading {
cancel()
leading = false
ctx, cancel = context.WithCancel(context.Background())
}
if !following {
following = true
go c.follower(ctx)
}
}
case <-c.stepdown:
if leading {
cancel()
c.r.LeadershipTransfer().Error()
}
case <-c.stop:
t.Stop()
cancel()
if leading {
c.r.LeadershipTransfer().Error()
}
close(c.done)
return nil
case <-t.C:
}
}
return nil
}
func (c *Cluster) leader(ctx context.Context) {
defer func() {
handleCrash()
}()
c.Callbacks.Leader(ctx, c.kv)
}
func (c *Cluster) follower(ctx context.Context) {
defer func() {
handleCrash()
}()
c.Callbacks.Follower(ctx, c.kv)
}
var reallyCrash = false
func handleCrash() {
if r := recover(); r != nil {
logPanic(r)
if reallyCrash {
panic(r)
}
}
}
// logPanic logs the caller tree when a panic occurs (except in the special case of http.ErrAbortHandler).
func logPanic(r interface{}) {
if r == http.ErrAbortHandler {
// honor the http.ErrAbortHandler sentinel panic value:
// ErrAbortHandler is a sentinel panic value to abort a handler.
// While any panic from ServeHTTP aborts the response to the client,
// panicking with ErrAbortHandler also suppresses logging of a stack trace to the server's error log.
return
}
// Same as stdlib http server code. Manually allocate stack trace buffer size
// to prevent excessively large logs
const size = 64 << 10
stacktrace := make([]byte, size)
stacktrace = stacktrace[:runtime.Stack(stacktrace, false)]
if _, ok := r.(string); ok {
klog.Errorf("Observed a panic: %s\n%s", r, stacktrace)
} else {
klog.Errorf("Observed a panic: %#v (%v)\n%s", r, r, stacktrace)
}
}
func (c *Cluster) Stepdown() {
c.stepdown <- struct{}{}
}
func (c *Cluster) Stop() {
close(c.stop)
<-c.done
s := make(chan struct{})
go func() {
c.r.Shutdown().Error()
close(s)
}()
select {
case <-s:
case <-time.After(time.Second * 5):
}
}

89
pkg/raft/tls/stream.go Normal file
View file

@ -0,0 +1,89 @@
package tls
import (
"crypto/tls"
"crypto/x509"
"io/ioutil"
"net"
"os"
"time"
"git.giftfish.de/ston1th/vipman/pkg/config"
"github.com/hashicorp/raft"
)
var (
ciphers = []uint16{
tls.TLS_CHACHA20_POLY1305_SHA256,
}
curves = []tls.CurveID{
tls.X25519,
}
)
func loadCertPool(file string) (p *x509.CertPool, err error) {
if file == "" {
return
}
f, err := os.Open(file)
if err != nil {
return
}
defer f.Close()
buf, err := ioutil.ReadAll(f)
if err != nil {
return
}
certs, err := x509.ParseCertificates(buf)
if err != nil {
return
}
p = x509.NewCertPool()
for _, c := range certs {
p.AddCert(c)
}
return
}
type stream struct {
net.Listener
insecure bool
ca *x509.CertPool
}
func NewStream(cfg *config.Raft) (s *stream, err error) {
cert, err := tls.LoadX509KeyPair(cfg.TLS.Cert, cfg.TLS.Key)
if err != nil {
return
}
ca, err := loadCertPool(cfg.TLS.CA)
if err != nil {
return
}
l, err := tls.Listen("tcp", cfg.Address, &tls.Config{
MinVersion: tls.VersionTLS13,
CipherSuites: ciphers,
CurvePreferences: curves,
PreferServerCipherSuites: true,
Certificates: []tls.Certificate{cert},
})
if err != nil {
return
}
s = &stream{l, cfg.TLS.Insecure, ca}
return
}
func (s *stream) Dial(address raft.ServerAddress, timeout time.Duration) (net.Conn, error) {
return tls.DialWithDialer(
&net.Dialer{Timeout: timeout},
"tcp", string(address),
&tls.Config{
MinVersion: tls.VersionTLS13,
CipherSuites: ciphers,
CurvePreferences: curves,
InsecureSkipVerify: s.insecure,
RootCAs: s.ca,
},
)
}