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
57 changes: 23 additions & 34 deletions go/raft/raft.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@ import (
"fmt"
"math/rand"
"net"
"strings"
"strconv"
"sync"
"sync/atomic"
"time"
Expand Down Expand Up @@ -93,7 +93,7 @@ func FatalRaftError(err error) error {
return err
}

func computeLeaderURI() (uri string, err error) {
func computeLeaderURI(raftAdvertise string) (uri string, err error) {
if config.Config.HTTPAdvertise != "" {
// Explicitly given
return config.Config.HTTPAdvertise, nil
Expand All @@ -104,14 +104,16 @@ func computeLeaderURI() (uri string, err error) {
scheme = "https"
}

hostname := strings.Split(config.Config.RaftAdvertise, ":")[0]
listenTokens := strings.Split(config.Config.ListenAddress, ":")
if len(listenTokens) < 2 {
hostname, _, err := net.SplitHostPort(raftAdvertise)
if err != nil {
hostname = raftAdvertise
}
_, port, err := net.SplitHostPort(config.Config.ListenAddress)
if err != nil {
return uri, fmt.Errorf("computeLeaderURI: cannot determine listen port out of config.Config.ListenAddress: %+v", config.Config.ListenAddress)
}
port := listenTokens[1]

uri = fmt.Sprintf("%s://%s:%s", scheme, hostname, port)
uri = fmt.Sprintf("%s://%s", scheme, net.JoinHostPort(hostname, port))
return uri, nil
}

Expand Down Expand Up @@ -146,7 +148,7 @@ func Setup(applier CommandApplier, snapshotCreatorApplier SnapshotCreatorApplier
return log.Errorf("failed to open raft store: %s", err.Error())
}

thisLeaderURI, err = computeLeaderURI()
thisLeaderURI, err = computeLeaderURI(raftAdvertise)
if err != nil {
return FatalRaftError(err)
}
Expand Down Expand Up @@ -175,37 +177,24 @@ func getRaft() *raft.Raft {
return store.raft
}

func normalizeRaftHostnameIP(host string) (string, error) {
if ip := net.ParseIP(host); ip != nil {
// this is a valid IP address.
return host, nil
}
ips, err := net.LookupIP(host)
if err != nil {
// resolve failed. But we don't want to fail the entire operation for that
_ = log.Errore(err)
return host, nil
}
// resolve success!
for _, ip := range ips {
return ip.String(), nil
}
return host, fmt.Errorf("%+v resolved but no IP found", host)
}

// normalizeRaftNode attempts to make sure there's a port to the given node.
// It consults the DefaultRaftPort when there isn't
// It consults the DefaultRaftPort when there isn't. The host is kept as
// given (never resolved to an IP), since it becomes this node's persistent
// raft ServerID.
func normalizeRaftNode(node string) (string, error) {
hostPort := strings.Split(node, ":")
host, err := normalizeRaftHostnameIP(hostPort[0])
host, port, err := net.SplitHostPort(node)
if err != nil {
return host, err
host, port = node, ""
}
if net.ParseIP(host) == nil {
if _, err := net.LookupHost(host); err != nil {
_ = log.Errore(err)
}
}
if len(hostPort) > 1 {
return fmt.Sprintf("%s:%s", host, hostPort[1]), nil
if port != "" {
return net.JoinHostPort(host, port), nil
} else if config.Config.DefaultRaftPort != 0 {
// No port specified, add one
return fmt.Sprintf("%s:%d", host, config.Config.DefaultRaftPort), nil
return net.JoinHostPort(host, strconv.Itoa(config.Config.DefaultRaftPort)), nil
} else {
return host, nil
}
Expand Down
117 changes: 117 additions & 0 deletions go/raft/raft_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,117 @@
package orcraft

import (
"testing"

"github.com/proxysql/orchestrator/go/config"
)

func TestNormalizeRaftNodePreservesHostnameWithPort(t *testing.T) {
node, err := normalizeRaftNode("orchestrator-0.orchestrator.default.svc.cluster.local:10008")
if err != nil {
t.Fatalf("normalizeRaftNode returned error: %+v", err)
}
if node != "orchestrator-0.orchestrator.default.svc.cluster.local:10008" {
t.Errorf("expected hostname:port to be preserved unresolved, got %q", node)
}
}

func TestNormalizeRaftNodePassesThroughLiteralIP(t *testing.T) {
node, err := normalizeRaftNode("192.168.1.10:10008")
if err != nil {
t.Fatalf("normalizeRaftNode returned error: %+v", err)
}
if node != "192.168.1.10:10008" {
t.Errorf("expected literal IP to pass through unchanged, got %q", node)
}
}

func withConfig(t *testing.T, httpAdvertise, listenAddress string, useSSL bool, fn func()) {
t.Helper()
origHTTPAdvertise := config.Config.HTTPAdvertise
origListenAddress := config.Config.ListenAddress
origUseSSL := config.Config.UseSSL
config.Config.HTTPAdvertise = httpAdvertise
config.Config.ListenAddress = listenAddress
config.Config.UseSSL = useSSL
defer func() {
config.Config.HTTPAdvertise = origHTTPAdvertise
config.Config.ListenAddress = origListenAddress
config.Config.UseSSL = origUseSSL
}()
fn()
}

func TestComputeLeaderURIWithIPv6Advertise(t *testing.T) {
withConfig(t, "", "0.0.0.0:3000", false, func() {
uri, err := computeLeaderURI("[::1]:10008")
if err != nil {
t.Fatalf("computeLeaderURI returned error: %+v", err)
}
if uri != "http://[::1]:3000" {
t.Errorf("expected %q, got %q", "http://[::1]:3000", uri)
}
})
}

func TestComputeLeaderURIWithHostname(t *testing.T) {
withConfig(t, "", "0.0.0.0:3000", false, func() {
uri, err := computeLeaderURI("orchestrator-0.svc.cluster.local:10008")
if err != nil {
t.Fatalf("computeLeaderURI returned error: %+v", err)
}
if uri != "http://orchestrator-0.svc.cluster.local:3000" {
t.Errorf("expected %q, got %q", "http://orchestrator-0.svc.cluster.local:3000", uri)
}
})
}

func TestComputeLeaderURIPrefersExplicitHTTPAdvertise(t *testing.T) {
withConfig(t, "https://explicit:9999", "0.0.0.0:3000", false, func() {
uri, err := computeLeaderURI("[::1]:10008")
if err != nil {
t.Fatalf("computeLeaderURI returned error: %+v", err)
}
if uri != "https://explicit:9999" {
t.Errorf("expected %q, got %q", "https://explicit:9999", uri)
}
})
}

func TestNormalizeRaftNodePassesThroughLiteralIPv6(t *testing.T) {
node, err := normalizeRaftNode("[::1]:10008")
if err != nil {
t.Fatalf("normalizeRaftNode returned error: %+v", err)
}
if node != "[::1]:10008" {
t.Errorf("expected literal IPv6 to pass through unchanged, got %q", node)
}
}

func TestNormalizeRaftNodeAddsDefaultPortToIPv6(t *testing.T) {
originalPort := config.Config.DefaultRaftPort
config.Config.DefaultRaftPort = 10008
defer func() { config.Config.DefaultRaftPort = originalPort }()

node, err := normalizeRaftNode("::1")
if err != nil {
t.Fatalf("normalizeRaftNode returned error: %+v", err)
}
if node != "[::1]:10008" {
t.Errorf("expected default port to be appended in bracketed form, got %q", node)
}
}

func TestNormalizeRaftNodeAddsDefaultPort(t *testing.T) {
originalPort := config.Config.DefaultRaftPort
config.Config.DefaultRaftPort = 10008
defer func() { config.Config.DefaultRaftPort = originalPort }()

node, err := normalizeRaftNode("orchestrator-0.orchestrator.default.svc.cluster.local")
if err != nil {
t.Fatalf("normalizeRaftNode returned error: %+v", err)
}
if node != "orchestrator-0.orchestrator.default.svc.cluster.local:10008" {
t.Errorf("expected default port to be appended, got %q", node)
}
}
84 changes: 84 additions & 0 deletions go/raft/store_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,84 @@
package orcraft

import (
"net"
"strconv"
"testing"
"time"

"github.com/hashicorp/raft"
)

// freeLocalPort asks the OS for a free TCP port on 127.0.0.1.
func freeLocalPort(t *testing.T) int {
t.Helper()
l, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("failed to allocate a free port: %+v", err)
}
port := l.Addr().(*net.TCPAddr).Port
if err := l.Close(); err != nil {
t.Fatalf("failed to release listener: %+v", err)
}
return port
}

// assertSingleNodeIdentity bootstraps a single-node store advertising as
// advertiseAddr and checks that raft's persisted configuration keys the node
// by that exact address, rather than some resolved variant of it.
func assertSingleNodeIdentity(t *testing.T, advertiseAddr string) {
t.Helper()
raftDir := t.TempDir()

raftBind, err := normalizeRaftNode(advertiseAddr)
if err != nil {
t.Fatalf("normalizeRaftNode failed: %+v", err)
}
raftAdvertise, err := normalizeRaftNode(advertiseAddr)
if err != nil {
t.Fatalf("normalizeRaftNode failed: %+v", err)
}

s := NewStore(raftDir, raftBind, raftAdvertise, nil, nil)
if err := s.Open(nil); err != nil {
t.Fatalf("Store.Open failed: %+v", err)
}
defer s.raft.Shutdown()

// Wait for the single-node cluster to elect itself leader.
deadline := time.Now().Add(5 * time.Second)
for time.Now().Before(deadline) {
if s.raft.State() == raft.Leader {
break
}
time.Sleep(50 * time.Millisecond)
}
if s.raft.State() != raft.Leader {
t.Fatalf("single-node raft store never became leader (state=%v)", s.raft.State())
}

future := s.raft.GetConfiguration()
if err := future.Error(); err != nil {
t.Fatalf("GetConfiguration failed: %+v", err)
}
servers := future.Configuration().Servers
if len(servers) != 1 {
t.Fatalf("expected exactly 1 server in configuration, got %d: %+v", len(servers), servers)
}

server := servers[0]
if string(server.ID) != advertiseAddr {
t.Errorf("expected raft ServerID to be %q, got %q (identity must not be a resolved, point-in-time address)", advertiseAddr, string(server.ID))
}
if string(server.Address) != advertiseAddr {
t.Errorf("expected raft ServerAddress to be %q, got %q", advertiseAddr, string(server.Address))
}
}

func TestStoreOpenKeepsHostnameAsServerIdentity(t *testing.T) {
assertSingleNodeIdentity(t, "localhost:"+strconv.Itoa(freeLocalPort(t)))
}

func TestStoreOpenKeepsIPv6AsServerIdentity(t *testing.T) {
assertSingleNodeIdentity(t, net.JoinHostPort("::1", strconv.Itoa(freeLocalPort(t))))
}