diff --git a/go/raft/raft.go b/go/raft/raft.go index 617bea2e6..bfa232ab4 100644 --- a/go/raft/raft.go +++ b/go/raft/raft.go @@ -21,7 +21,7 @@ import ( "fmt" "math/rand" "net" - "strings" + "strconv" "sync" "sync/atomic" "time" @@ -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 @@ -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 } @@ -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) } @@ -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 } diff --git a/go/raft/raft_test.go b/go/raft/raft_test.go new file mode 100644 index 000000000..5b3064b33 --- /dev/null +++ b/go/raft/raft_test.go @@ -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) + } +} diff --git a/go/raft/store_test.go b/go/raft/store_test.go new file mode 100644 index 000000000..2c6e03a88 --- /dev/null +++ b/go/raft/store_test.go @@ -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)))) +}