diff --git a/.github/workflows/lint-test.yaml b/.github/workflows/lint-test.yaml index 37ac0f0..9769467 100644 --- a/.github/workflows/lint-test.yaml +++ b/.github/workflows/lint-test.yaml @@ -8,7 +8,7 @@ on: # yamllint disable-line rule:truthy pull_request: branches: ["*"] env: - GO_VERSION: "~1.23.0" + GO_VERSION: "~1.25.0" jobs: go-lint: name: "Lint Go" diff --git a/balancer.go b/balancer.go index d48a961..cca0cd2 100644 --- a/balancer.go +++ b/balancer.go @@ -146,13 +146,15 @@ func (b *builder) Name() string { return BalancerName } func (b *builder) Build(cc balancer.ClientConn, _ balancer.BuildOptions) balancer.Balancer { bal := &ringBalancer{ - cc: cc, - subConns: resolver.NewAddressMap(), - scStates: make(map[balancer.SubConn]connectivity.State), - csEvltr: &balancer.ConnectivityStateEvaluator{}, - state: connectivity.Connecting, - hasher: b.hashfn, - picker: base.NewErrPicker(balancer.ErrNoSubConnAvailable), + cc: cc, + subConns: resolver.NewAddressMapV2[any](), + scStates: make(map[balancer.SubConn]connectivity.State), + scKeys: make(map[balancer.SubConn]string), + ringMembers: make(map[balancer.SubConn]struct{}), + csEvltr: &balancer.ConnectivityStateEvaluator{}, + state: connectivity.Connecting, + hasher: b.hashfn, + picker: base.NewErrPicker(balancer.ErrNoSubConnAvailable), } return bal @@ -186,8 +188,12 @@ type ringBalancer struct { cc balancer.ClientConn picker balancer.Picker csEvltr *balancer.ConnectivityStateEvaluator - subConns *resolver.AddressMap + subConns *resolver.AddressMapV2[any] scStates map[balancer.SubConn]connectivity.State + // scKeys holds the hashring key of every SubConn the resolver gave us. + scKeys map[balancer.SubConn]string + // ringMembers is the set of SubConns on the hashring: only READY ones. + ringMembers map[balancer.SubConn]struct{} config *BalancerConfig hashring *hashring.Ring @@ -235,8 +241,18 @@ func (b *ringBalancer) UpdateClientConnState(s balancer.ClientConnState) error { svcConfig := s.BalancerConfig.(*BalancerConfig) if b.config == nil || svcConfig.ReplicationFactor != b.config.ReplicationFactor { b.hashring = hashring.MustNew(b.hasher, svcConfig.ReplicationFactor) - b.config = svcConfig + // The new ring starts empty: put every READY SubConn back on it. + b.ringMembers = make(map[balancer.SubConn]struct{}) + for sc, st := range b.scStates { + if st != connectivity.Ready { + continue + } + if err := b.addToRing(sc); err != nil { + return err + } + } } + b.config = svcConfig } // if there's no hashring yet, the balancer hasn't yet parsed an initial @@ -252,12 +268,16 @@ func (b *ringBalancer) UpdateClientConnState(s balancer.ClientConnState) error { // if any new targets have been added, they are added to the hashring, and // any that have been removed since the last update are removed from the // hashring. - addrsSet := resolver.NewAddressMap() + addrsSet := resolver.NewAddressMapV2[any]() for _, addr := range s.ResolverState.Addresses { addrsSet.Set(addr, nil) if _, ok := b.subConns.Get(addr); !ok { // addr is addr new address (not existing in b.subConns). + // NewSubConn is deprecated only to warn that a SubConn will soon + // hold a single address. We already pass exactly one address, and + // grpc offers no replacement API yet. + //nolint:staticcheck // SA1019: single-address usage is the future-proof form sc, err := b.cc.NewSubConn([]resolver.Address{addr}, balancer.NewSubConnOptions{HealthCheckEnabled: false}) if err != nil { logger.Warningf("base.baseBalancer: failed to create new SubConn: %v", err) @@ -266,15 +286,9 @@ func (b *ringBalancer) UpdateClientConnState(s balancer.ClientConnState) error { b.subConns.Set(addr, sc) b.scStates[sc] = connectivity.Idle + b.scKeys[sc] = addr.ServerName + addr.Addr b.csEvltr.RecordTransition(connectivity.Shutdown, connectivity.Idle) sc.Connect() - - if err := b.hashring.Add(subConnMember{ - SubConn: sc, - key: addr.ServerName + addr.Addr, - }); err != nil { - return fmt.Errorf("couldn't add to hashring") - } } } @@ -283,15 +297,12 @@ func (b *ringBalancer) UpdateClientConnState(s balancer.ClientConnState) error { sc := sci.(balancer.SubConn) // addr was removed by resolver. if _, ok := addrsSet.Get(addr); !ok { - b.cc.RemoveSubConn(sc) + sc.Shutdown() b.subConns.Delete(addr) // Keep the state of this sc in b.scStates until sc's state becomes Shutdown. // The entry will be deleted in UpdateSubConnState. - if err := b.hashring.Remove(subConnMember{ - SubConn: sc, - key: addr.ServerName + addr.Addr, - }); err != nil { - return fmt.Errorf("couldn't add to hashring") + if err := b.removeFromRing(sc); err != nil { + return err } } } @@ -313,19 +324,52 @@ func (b *ringBalancer) UpdateClientConnState(s balancer.ClientConnState) error { return balancer.ErrBadResolverState } - // If the overall connection state is not in transient failure, we return - // addr new picker with addr reference to the hashring (otherwise an error picker) + b.regeneratePicker() + b.cc.UpdateState(balancer.State{ConnectivityState: b.state, Picker: b.picker}) + + return nil +} + +// regeneratePicker installs an error picker while the balancer is in +// TRANSIENT_FAILURE and a hashring picker otherwise. +func (b *ringBalancer) regeneratePicker() { if b.state == connectivity.TransientFailure { b.picker = base.NewErrPicker(errors.Join(b.connErr, b.resolverErr)) - } else { - b.picker = &picker{ - hashring: b.hashring, - spread: b.config.Spread, - } + return } - // update the ClientConn with the current hashring picker picker - b.cc.UpdateState(balancer.State{ConnectivityState: b.state, Picker: b.picker}) + b.picker = &picker{ + hashring: b.hashring, + spread: b.config.Spread, + } +} + +// addToRing places sc on the hashring if it is not already there. +func (b *ringBalancer) addToRing(sc balancer.SubConn) error { + if _, ok := b.ringMembers[sc]; ok { + return nil + } + + if err := b.hashring.Add(subConnMember{SubConn: sc, key: b.scKeys[sc]}); err != nil { + return fmt.Errorf("couldn't add to hashring: %w", err) + } + + b.ringMembers[sc] = struct{}{} + + return nil +} + +// removeFromRing takes sc off the hashring if it is there. +func (b *ringBalancer) removeFromRing(sc balancer.SubConn) error { + if _, ok := b.ringMembers[sc]; !ok { + return nil + } + + if err := b.hashring.Remove(subConnMember{SubConn: sc, key: b.scKeys[sc]}); err != nil { + return fmt.Errorf("couldn't remove from hashring: %w", err) + } + + delete(b.ringMembers, sc) return nil } @@ -362,6 +406,19 @@ func (b *ringBalancer) UpdateSubConnState(sc balancer.SubConn, state balancer.Su b.scStates[sc] = s + // A backend only receives traffic while its connection is READY. Keys + // hashed to a backend that is connecting or failed move to the closest + // ready member instead of waiting on a connection that may never come. + var ringErr error + if s == connectivity.Ready { + ringErr = b.addToRing(sc) + } else { + ringErr = b.removeFromRing(sc) + } + if ringErr != nil { + logger.Warningf("consistent-hashring: %v", ringErr) + } + switch s { case connectivity.Idle: sc.Connect() @@ -369,12 +426,14 @@ func (b *ringBalancer) UpdateSubConnState(sc balancer.SubConn, state balancer.Su // When an address was removed by resolver, b called RemoveSubConn but // kept the sc's state in scStates. Remove state for this sc here. delete(b.scStates, sc) + delete(b.scKeys, sc) case connectivity.TransientFailure: // Save error to be reported via picker. b.connErr = state.ConnectionError } b.state = b.csEvltr.RecordTransition(oldS, s) + b.regeneratePicker() b.cc.UpdateState(balancer.State{ConnectivityState: b.state, Picker: b.picker}) } @@ -401,10 +460,10 @@ var _ balancer.Picker = (*picker)(nil) // The value stored in CtxKey is hashed into the hashring, and the resulting // subconnection is used. // -// There is no fallback behavior if the subconnection is unavailable; this -// prevents the request from going to a node that doesn't expect to receive it. -// As long as you are using a resolver that removes connections from the list -// when they are observably unavailable, this is a non-issue. +// The hashring only contains READY subconnections, so a key whose closest +// backend is down is served by the next closest ready backend. When no +// backend is ready the RPC is queued until one becomes ready (or the +// balancer reports TRANSIENT_FAILURE and installs an error picker). // // Spread can be increased to be robust against single node availability // problems. If spread is greater than 1, a random selection is made from the @@ -412,14 +471,20 @@ var _ balancer.Picker = (*picker)(nil) func (p *picker) Pick(info balancer.PickInfo) (balancer.PickResult, error) { key := info.Ctx.Value(CtxKey).([]byte) + // FindN only fails with hashring.ErrNotEnoughMembers. members, err := p.hashring.FindN(key, p.spread) if err != nil { - return balancer.PickResult{}, err + // Fewer ready backends than the configured spread: use those that are. + members, err = p.hashring.FindN(key, 1) + if err != nil { + // No ready backends at all: queue the RPC until one is ready. + return balancer.PickResult{}, balancer.ErrNoSubConnAvailable + } } index := 0 - if p.spread > 1 { - index = intn(p.spread) + if len(members) > 1 { + index = intn(uint8(len(members))) } chosen := members[index].(subConnMember) diff --git a/balancer_readiness_test.go b/balancer_readiness_test.go new file mode 100644 index 0000000..c11fc66 --- /dev/null +++ b/balancer_readiness_test.go @@ -0,0 +1,304 @@ +package consistent + +import ( + "context" + "fmt" + "net" + "testing" + "time" + + "github.com/cespare/xxhash/v2" + "github.com/stretchr/testify/require" + "google.golang.org/grpc" + "google.golang.org/grpc/balancer" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/connectivity" + "google.golang.org/grpc/credentials/insecure" + "google.golang.org/grpc/health" + healthpb "google.golang.org/grpc/health/grpc_health_v1" + "google.golang.org/grpc/resolver" + "google.golang.org/grpc/resolver/manual" + "google.golang.org/grpc/status" + + "github.com/authzed/consistent/hashring" +) + +// keyHashingTo returns a key that the ring places on the member with the +// given key, so a test can target a specific backend. +func keyHashingTo(t *testing.T, ring *hashring.Ring, memberKey string) []byte { + t.Helper() + for i := 0; i < 10000; i++ { + key := []byte(fmt.Sprintf("key-%d", i)) + members, err := ring.FindN(key, 1) + require.NoError(t, err) + if members[0].Key() == memberKey { + return key + } + } + t.Fatalf("no key hashes to %s", memberKey) + return nil +} + +// TestPickerSkipsSubConnsThatAreNotReady drives the balancer directly: two +// backends come up, then one fails. Keys that hashed to the failed backend +// must be picked on the remaining ready one instead of a dead SubConn. +func TestPickerSkipsSubConnsThatAreNotReady(t *testing.T) { + cc := newFakeClientConn() + // Drain state updates so the balancer never blocks on the fake conn. + go func() { + for range cc.stateCh { + } + }() + + b := NewBuilder(xxhash.Sum64).Build(cc, balancer.BuildOptions{}).(*ringBalancer) + addrs := []resolver.Address{{ServerName: "t", Addr: "1"}, {ServerName: "t", Addr: "2"}} + require.NoError(t, b.UpdateClientConnState(balancer.ClientConnState{ + ResolverState: resolver.State{Addresses: addrs}, + BalancerConfig: &BalancerConfig{ReplicationFactor: 100, Spread: 1}, + })) + + subConnFor := func(addr resolver.Address) balancer.SubConn { + sc, ok := b.subConns.Get(addr) + require.True(t, ok) + return sc.(balancer.SubConn) + } + sc1, sc2 := subConnFor(addrs[0]), subConnFor(addrs[1]) + for _, sc := range []balancer.SubConn{sc1, sc2} { + b.UpdateSubConnState(sc, balancer.SubConnState{ConnectivityState: connectivity.Connecting}) + b.UpdateSubConnState(sc, balancer.SubConnState{ConnectivityState: connectivity.Ready}) + } + + // Build a reference ring to learn which keys land on backend 2. + ref := hashring.MustNew(xxhash.Sum64, 100) + for _, a := range addrs { + require.NoError(t, ref.Add(subConnMember{key: a.ServerName + a.Addr})) + } + keyFor2 := keyHashingTo(t, ref, "t2") + + pick := func(key []byte) balancer.SubConn { + res, err := b.picker.Pick(balancer.PickInfo{Ctx: context.WithValue(context.Background(), CtxKey, key)}) + require.NoError(t, err) + return res.SubConn + } + require.Same(t, sc2, pick(keyFor2), "sanity: key routes to backend 2 while it is ready") + + // Backend 2 dies. + b.UpdateSubConnState(sc2, balancer.SubConnState{ + ConnectivityState: connectivity.TransientFailure, + ConnectionError: fmt.Errorf("connection refused"), + }) + require.Same(t, sc1, pick(keyFor2), "key must move to the only ready backend") + + // Backend 2 comes back and is used again. + b.UpdateSubConnState(sc2, balancer.SubConnState{ConnectivityState: connectivity.Idle}) + b.UpdateSubConnState(sc2, balancer.SubConnState{ConnectivityState: connectivity.Connecting}) + b.UpdateSubConnState(sc2, balancer.SubConnState{ConnectivityState: connectivity.Ready}) + require.Same(t, sc2, pick(keyFor2), "key returns to backend 2 once it is ready") +} + +// TestRPCsAreNotStuckBehindAConnectingBackend reproduces a peer whose address +// is still resolvable but never completes a connection (a killed pod whose IP +// is still in the endpoint list). RPCs hashed to it must be served by the +// healthy backend instead of waiting for the dial to time out. +func TestRPCsAreNotStuckBehindAConnectingBackend(t *testing.T) { + // Healthy backend: a real gRPC server with the health service. + healthyLis, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + srv := grpc.NewServer() + healthpb.RegisterHealthServer(srv, health.NewServer()) + go func() { _ = srv.Serve(healthyLis) }() + t.Cleanup(srv.Stop) + + // Black hole: accepts TCP connections but never speaks HTTP/2, so the + // SubConn stays CONNECTING until gRPC's 20s connect timeout. + blackholeLis, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + t.Cleanup(func() { _ = blackholeLis.Close() }) + go func() { + for { + conn, err := blackholeLis.Accept() + if err != nil { + return + } + t.Cleanup(func() { _ = conn.Close() }) + } + }() + + balancer.Register(NewBuilder(xxhash.Sum64)) + addrs := []resolver.Address{{Addr: healthyLis.Addr().String()}, {Addr: blackholeLis.Addr().String()}} + rb := manual.NewBuilderWithScheme("readiness") + rb.InitialState(resolver.State{Addresses: addrs}) + + svcConfig, err := (&BalancerConfig{ReplicationFactor: 100, Spread: 1}).ServiceConfigJSON() + require.NoError(t, err) + conn, err := grpc.NewClient(rb.Scheme()+":///backends", + grpc.WithTransportCredentials(insecure.NewCredentials()), + grpc.WithResolvers(rb), + grpc.WithDefaultServiceConfig(svcConfig), + ) + require.NoError(t, err) + t.Cleanup(func() { _ = conn.Close() }) + client := healthpb.NewHealthClient(conn) + + ref := hashring.MustNew(xxhash.Sum64, 100) + for _, a := range addrs { + require.NoError(t, ref.Add(subConnMember{key: a.ServerName + a.Addr})) + } + keyForHealthy := keyHashingTo(t, ref, addrs[0].Addr) + keyForBlackhole := keyHashingTo(t, ref, addrs[1].Addr) + + check := func(key []byte, timeout time.Duration) error { + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + ctx = context.WithValue(ctx, CtxKey, key) + _, err := client.Check(ctx, &healthpb.HealthCheckRequest{}) + return err + } + + // Warm up: the healthy backend is reachable and connected. + require.NoError(t, check(keyForHealthy, 5*time.Second)) + + // A key hashed to the black-holed backend must still be answered promptly. + err = check(keyForBlackhole, 3*time.Second) + if st, ok := status.FromError(err); ok && st.Code() == codes.DeadlineExceeded { + t.Fatalf("RPC hung behind a CONNECTING backend instead of being served by the ready one: %v", err) + } + require.NoError(t, err) +} + +// readyBalancer returns a balancer whose SubConns for addrs are all READY, +// with a buffered fake ClientConn so state pushes can be counted. +func readyBalancer(t *testing.T, addrs ...resolver.Address) (*ringBalancer, *fakeClientConn) { + t.Helper() + cc := newFakeClientConn() + cc.stateCh = make(chan balancer.State, 64) + b := NewBuilder(xxhash.Sum64).Build(cc, balancer.BuildOptions{}).(*ringBalancer) + require.NoError(t, b.UpdateClientConnState(balancer.ClientConnState{ + ResolverState: resolver.State{Addresses: addrs}, + BalancerConfig: &BalancerConfig{ReplicationFactor: 100, Spread: 1}, + })) + for _, sci := range b.subConns.Values() { + sc := sci.(balancer.SubConn) + b.UpdateSubConnState(sc, balancer.SubConnState{ConnectivityState: connectivity.Connecting}) + b.UpdateSubConnState(sc, balancer.SubConnState{ConnectivityState: connectivity.Ready}) + } + return b, cc +} + +func ringKeys(b *ringBalancer) []string { + return keys(b.hashring.Members()) +} + +func TestConnectingSubConnLeavesTheRing(t *testing.T) { + addr := resolver.Address{ServerName: "t", Addr: "1"} + b, _ := readyBalancer(t, addr) + sci, _ := b.subConns.Get(addr) + sc := sci.(balancer.SubConn) + require.Equal(t, []string{"t1"}, ringKeys(b)) + + b.UpdateSubConnState(sc, balancer.SubConnState{ConnectivityState: connectivity.Connecting}) + require.Empty(t, ringKeys(b), "a reconnecting backend must not receive traffic") + require.Equal(t, connectivity.Connecting, b.state) +} + +// After TRANSIENT_FAILURE, IDLE and CONNECTING reports are ignored so the +// aggregate state does not flap, but an IDLE SubConn is still asked to +// reconnect. +func TestTransientFailureIgnoresReconnectTransitions(t *testing.T) { + addr := resolver.Address{ServerName: "t", Addr: "1"} + b, _ := readyBalancer(t, addr) + sci, _ := b.subConns.Get(addr) + sc := sci.(*fakeSubConn) + connectsSoFar := sc.connectCalls + + b.UpdateSubConnState(sc, balancer.SubConnState{ + ConnectivityState: connectivity.TransientFailure, + ConnectionError: fmt.Errorf("refused"), + }) + require.Equal(t, connectivity.TransientFailure, b.state) + require.Empty(t, ringKeys(b)) + + b.UpdateSubConnState(sc, balancer.SubConnState{ConnectivityState: connectivity.Connecting}) + require.Equal(t, connectivity.TransientFailure, b.scStates[sc]) + require.Equal(t, connectivity.TransientFailure, b.state) + require.Equal(t, connectsSoFar, sc.connectCalls, "CONNECTING must not trigger another Connect") + + b.UpdateSubConnState(sc, balancer.SubConnState{ConnectivityState: connectivity.Idle}) + require.Equal(t, connectivity.TransientFailure, b.scStates[sc]) + require.Equal(t, connectivity.TransientFailure, b.state) + require.Equal(t, connectsSoFar+1, sc.connectCalls, "IDLE must trigger exactly one Connect") + + b.UpdateSubConnState(sc, balancer.SubConnState{ConnectivityState: connectivity.Ready}) + require.Equal(t, connectivity.Ready, b.state) + require.Equal(t, []string{"t1"}, ringKeys(b)) +} + +func TestResolverErrorOnlyPushesStateInTransientFailure(t *testing.T) { + b, cc := readyBalancer(t, resolver.Address{ServerName: "t", Addr: "1"}) + for len(cc.stateCh) > 0 { + <-cc.stateCh + } + + b.ResolverError(fmt.Errorf("dns down")) + require.Equal(t, connectivity.Ready, b.state) + require.Empty(t, cc.stateCh, "a healthy balancer keeps its picker on resolver errors") + + // With no SubConns at all the error is surfaced through an error picker. + empty := NewBuilder(xxhash.Sum64).Build(newFakeClientConn(), balancer.BuildOptions{}).(*ringBalancer) + emptyCC := empty.cc.(*fakeClientConn) + emptyCC.stateCh = make(chan balancer.State, 1) + empty.ResolverError(fmt.Errorf("dns down")) + require.Equal(t, connectivity.TransientFailure, empty.state) + require.Len(t, emptyCC.stateCh, 1) + pushed := <-emptyCC.stateCh + require.Equal(t, connectivity.TransientFailure, pushed.ConnectivityState) + _, err := pushed.Picker.Pick(balancer.PickInfo{}) + require.EqualError(t, err, "dns down") +} + +func TestReplicationFactorChangeKeepsReadyMembers(t *testing.T) { + addrs := []resolver.Address{{ServerName: "t", Addr: "1"}, {ServerName: "t", Addr: "2"}} + b, _ := readyBalancer(t, addrs...) + before := b.hashring + + require.NoError(t, b.UpdateClientConnState(balancer.ClientConnState{ + ResolverState: resolver.State{Addresses: addrs}, + BalancerConfig: &BalancerConfig{ReplicationFactor: 7, Spread: 1}, + })) + require.NotSame(t, before, b.hashring, "a new replication factor builds a new ring") + require.ElementsMatch(t, []string{"t1", "t2"}, ringKeys(b)) + + // Same factor: the ring is kept as is. + current := b.hashring + require.NoError(t, b.UpdateClientConnState(balancer.ClientConnState{ + ResolverState: resolver.State{Addresses: addrs}, + BalancerConfig: &BalancerConfig{ReplicationFactor: 7, Spread: 2}, + })) + require.Same(t, current, b.hashring) + require.Equal(t, uint8(2), b.picker.(*picker).spread) +} + +func TestParseConfigDefaults(t *testing.T) { + bld := NewBuilder(xxhash.Sum64) + + cfg, err := bld.ParseConfig([]byte(`{}`)) + require.NoError(t, err) + require.Equal(t, &BalancerConfig{ReplicationFactor: DefaultReplicationFactor, Spread: DefaultSpread}, cfg) + + cfg, err = bld.ParseConfig([]byte(`{"replicationFactor": 3, "spread": 2}`)) + require.NoError(t, err) + require.Equal(t, &BalancerConfig{ReplicationFactor: 3, Spread: 2}, cfg) + + _, err = bld.ParseConfig([]byte(`not json`)) + require.Error(t, err) +} + +func TestIntnStaysInRange(t *testing.T) { + for _, n := range []uint8{1, 2, 3, 7} { + for i := 0; i < 2000; i++ { + v := intn(n) + require.GreaterOrEqual(t, v, 0) + require.Less(t, v, int(n)) + } + } +} diff --git a/balancer_test.go b/balancer_test.go index c36b79a..b692a34 100644 --- a/balancer_test.go +++ b/balancer_test.go @@ -22,10 +22,15 @@ import ( type fakeSubConn struct { balancer.SubConn - id string + id string + connectCalls int + shutdownCalls int } -func (fakeSubConn) Connect() {} +func (sc *fakeSubConn) Connect() { sc.connectCalls++ } + +// Shutdown is a no-op so the embedded nil SubConn is never dereferenced. +func (sc *fakeSubConn) Shutdown() { sc.shutdownCalls++ } func keys(members []hashring.Member) []string { keys := make([]string, 0, len(members)) @@ -39,6 +44,8 @@ func keys(members []hashring.Member) []string { // behavior itself, see `pkg/consistent` for tests of the hashring. func TestConsistentHashringPickerPick(t *testing.T) { // Override the intn function with one that uses a stable seed. + realIntn := intn + t.Cleanup(func() { intn = realIntn }) intn = func(n uint8) int { h := new(maphash.Hash) @@ -155,26 +162,24 @@ func TestConsistentHashringBalancerConfigServiceConfigJSON(t *testing.T) { } func TestConsistentHashringBalancerUpdateClientConnState(t *testing.T) { + // Each step applies a resolver update, then brings every SubConn the + // resolver produced to READY, and finally checks the installed picker. type balancerState struct { ConnectivityState connectivity.State err error memberKeys []string spread uint8 - replicationFactor uint16 } tests := []struct { - name string - s []balancer.ClientConnState - expectedStates []balancerState - expectedConnState connectivity.State - wantErr bool + name string + s []balancer.ClientConnState + expectedStates []balancerState + wantErr bool }{ { - name: "no hashring", - expectedStates: []balancerState{}, - expectedConnState: connectivity.TransientFailure, - wantErr: true, + name: "no hashring", + wantErr: true, }, { name: "configures hashring, no addresses", @@ -191,8 +196,7 @@ func TestConsistentHashringBalancerUpdateClientConnState(t *testing.T) { err: errors.Join(nil, fmt.Errorf("produced zero addresses")), }, }, - expectedConnState: connectivity.TransientFailure, - wantErr: true, + wantErr: true, }, { name: "configures hashring, 3 addresses", @@ -211,13 +215,11 @@ func TestConsistentHashringBalancerUpdateClientConnState(t *testing.T) { }}, expectedStates: []balancerState{ { - ConnectivityState: connectivity.Connecting, + ConnectivityState: connectivity.Ready, memberKeys: []string{"t1", "t2", "t3"}, - replicationFactor: 100, spread: 1, }, }, - expectedConnState: connectivity.Idle, }, { name: "existing hashring with 3 nodes, 1 removed", @@ -243,19 +245,50 @@ func TestConsistentHashringBalancerUpdateClientConnState(t *testing.T) { }}, expectedStates: []balancerState{ { - ConnectivityState: connectivity.Connecting, + ConnectivityState: connectivity.Ready, memberKeys: []string{"t1", "t2", "t3"}, - replicationFactor: 100, spread: 1, }, { - ConnectivityState: connectivity.Connecting, + ConnectivityState: connectivity.Ready, memberKeys: []string{"t1", "t2"}, - replicationFactor: 100, spread: 1, }, }, - expectedConnState: connectivity.Idle, + }, + { + name: "existing hashring with 3 nodes, 2 removed", + s: []balancer.ClientConnState{{ + ResolverState: resolver.State{ + Addresses: []resolver.Address{ + {ServerName: "t", Addr: "1"}, + {ServerName: "t", Addr: "2"}, + {ServerName: "t", Addr: "3"}, + }, + }, + BalancerConfig: &BalancerConfig{ + ReplicationFactor: 100, + Spread: 1, + }, + }, { + ResolverState: resolver.State{ + Addresses: []resolver.Address{ + {ServerName: "t", Addr: "2"}, + }, + }, + }}, + expectedStates: []balancerState{ + { + ConnectivityState: connectivity.Ready, + memberKeys: []string{"t1", "t2", "t3"}, + spread: 1, + }, + { + ConnectivityState: connectivity.Ready, + memberKeys: []string{"t2"}, + spread: 1, + }, + }, }, { name: "existing hashring with 3 nodes, 1 added", @@ -283,19 +316,16 @@ func TestConsistentHashringBalancerUpdateClientConnState(t *testing.T) { }}, expectedStates: []balancerState{ { - ConnectivityState: connectivity.Connecting, + ConnectivityState: connectivity.Ready, memberKeys: []string{"t1", "t2", "t3"}, - replicationFactor: 100, spread: 1, }, { - ConnectivityState: connectivity.Connecting, + ConnectivityState: connectivity.Ready, memberKeys: []string{"t1", "t2", "t3", "t4"}, - replicationFactor: 100, spread: 1, }, }, - expectedConnState: connectivity.Idle, }, { name: "existing hashring with 3 nodes, 1 replaced", @@ -322,67 +352,64 @@ func TestConsistentHashringBalancerUpdateClientConnState(t *testing.T) { }}, expectedStates: []balancerState{ { - ConnectivityState: connectivity.Connecting, + ConnectivityState: connectivity.Ready, memberKeys: []string{"t1", "t2", "t3"}, - replicationFactor: 100, spread: 1, }, { - ConnectivityState: connectivity.Connecting, + ConnectivityState: connectivity.Ready, memberKeys: []string{"t1", "t2", "t4"}, - replicationFactor: 100, spread: 1, }, }, - expectedConnState: connectivity.Idle, }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { b := NewBuilder(xxhash.Sum64) cc := newFakeClientConn() - bb := b.Build(cc, balancer.BuildOptions{}) - cb := bb.(*ringBalancer) - - tt := tt - - done := make(chan struct{}) + cb := b.Build(cc, balancer.BuildOptions{}).(*ringBalancer) + // The balancer pushes a state on every change; the assertions below + // look at the balancer itself, so the pushes only need draining. go func() { - i := 0 - - if len(tt.expectedStates) == 0 { - done <- struct{}{} - return + for range cc.stateCh { } + }() + defer close(cc.stateCh) - for { - s := <-cc.stateCh - expected := tt.expectedStates[i] - require.Equal(t, expected.ConnectivityState, s.ConnectivityState) + if len(tt.s) == 0 { + err := cb.UpdateClientConnState(balancer.ClientConnState{}) + require.Equal(t, tt.wantErr, err != nil) + require.Equal(t, base.NewErrPicker(nil), cb.picker) + return + } - if expected.err != nil { - require.Equal(t, base.NewErrPicker(expected.err), s.Picker) - } else { - p := s.Picker.(*picker) - require.Equal(t, expected.spread, p.spread) - require.ElementsMatch(t, expected.memberKeys, keys(p.hashring.Members())) - } + for i, state := range tt.s { + if err := cb.UpdateClientConnState(state); (err != nil) != tt.wantErr { + t.Errorf("UpdateClientConnState() error = %v, wantErr %v", err, tt.wantErr) + } - i++ - done <- struct{}{} + for _, sci := range cb.subConns.Values() { + sc := sci.(balancer.SubConn) + if cb.scStates[sc] == connectivity.Ready { + continue + } + cb.UpdateSubConnState(sc, balancer.SubConnState{ConnectivityState: connectivity.Connecting}) + cb.UpdateSubConnState(sc, balancer.SubConnState{ConnectivityState: connectivity.Ready}) } - }() - for _, state := range tt.s { - if err := cb.UpdateClientConnState(state); (err != nil) != tt.wantErr { - t.Errorf("UpdateClientConnState() error = %v, wantErr %v", err, tt.wantErr) + expected := tt.expectedStates[i] + require.Equal(t, expected.ConnectivityState, cb.state) + if expected.err != nil { + require.Equal(t, base.NewErrPicker(expected.err), cb.picker) + continue } - <-done + p := cb.picker.(*picker) + require.Equal(t, expected.spread, p.spread) + require.ElementsMatch(t, expected.memberKeys, keys(p.hashring.Members())) } - - require.Equal(t, tt.expectedConnState, cb.csEvltr.CurrentState()) }) } } diff --git a/go.mod b/go.mod index 10c29a1..abf1ac5 100644 --- a/go.mod +++ b/go.mod @@ -1,22 +1,18 @@ module github.com/authzed/consistent -go 1.23 +go 1.25.0 require ( - github.com/cespare/xxhash/v2 v2.2.0 - github.com/stretchr/testify v1.8.4 - golang.org/x/exp v0.0.0-20230801115018-d63ba01acd4b - google.golang.org/grpc v1.56.2 + github.com/cespare/xxhash/v2 v2.3.0 + github.com/stretchr/testify v1.12.1 + google.golang.org/grpc v1.83.2 ) require ( - github.com/davecgh/go-spew v1.1.1 // indirect - github.com/golang/protobuf v1.5.3 // indirect - github.com/kr/pretty v0.3.1 // indirect - github.com/pmezard/go-difflib v1.0.0 // indirect - github.com/rogpeppe/go-internal v1.10.0 // indirect - golang.org/x/sys v0.9.0 // indirect - google.golang.org/protobuf v1.31.0 // indirect - gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c // indirect - gopkg.in/yaml.v3 v3.0.1 // indirect + go.yaml.in/yaml/v3 v3.0.5 // indirect + golang.org/x/net v0.58.0 // indirect + golang.org/x/sys v0.47.0 // indirect + golang.org/x/text v0.41.0 // indirect + google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect + google.golang.org/protobuf v1.36.11 // indirect ) diff --git a/go.sum b/go.sum index 3c875e5..44f36f5 100644 --- a/go.sum +++ b/go.sum @@ -1,42 +1,42 @@ -github.com/cespare/xxhash/v2 v2.2.0 h1:DC2CZ1Ep5Y4k3ZQ899DldepgrayRUGE6BBZ/cd9Cj44= -github.com/cespare/xxhash/v2 v2.2.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= -github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E= -github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= -github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= -github.com/golang/protobuf v1.5.0/go.mod h1:FsONVRAS9T7sI+LIUmWTfcYkHO4aIWwzhcaSAoJOfIk= -github.com/golang/protobuf v1.5.3 h1:KhyjKVUg7Usr/dYsdSqoFveMYd5ko72D+zANwlG1mmg= -github.com/golang/protobuf v1.5.3/go.mod h1:XVQd3VNwM+JqD3oG2Ue2ip4fOMUkwXdXDdiuN0vRsmY= -github.com/google/go-cmp v0.5.5/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE= -github.com/google/go-cmp v0.5.9 h1:O2Tfq5qg4qc4AmwVlvv0oLiVAGB7enBSJ2x2DqQFi38= -github.com/google/go-cmp v0.5.9/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= -github.com/kr/pretty v0.2.1/go.mod h1:ipq/a2n7PKx3OHsz4KJII5eveXtPO4qwEXGdVfWzfnI= -github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= -github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= -github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ= -github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI= -github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= -github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= -github.com/pkg/diff v0.0.0-20210226163009-20ebb0f2a09e/go.mod h1:pJLUxLENpZxwdsKMEsNbx1VGcRFpLqf3715MtcvvzbA= -github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= -github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= -github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs= -github.com/rogpeppe/go-internal v1.10.0 h1:TMyTOH3F/DB16zRVcYyreMH6GnZZrwQVAoYjRBZyWFQ= -github.com/rogpeppe/go-internal v1.10.0/go.mod h1:UQnix2H7Ngw/k4C5ijL5+65zddjncjaFoBhdsK/akog= -github.com/stretchr/testify v1.8.4 h1:CcVxjf3Q8PM0mHUKJCdn+eZZtm5yQwehR5yeSVQQcUk= -github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo= -golang.org/x/exp v0.0.0-20230801115018-d63ba01acd4b h1:r+vk0EmXNmekl0S0BascoeeoHk/L7wmaW2QF90K+kYI= -golang.org/x/exp v0.0.0-20230801115018-d63ba01acd4b/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc= -golang.org/x/sys v0.9.0 h1:KS/R3tvhPqvJvwcKfnBHJwwthS11LRhmM5D59eEXa0s= -golang.org/x/sys v0.9.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= -google.golang.org/grpc v1.56.2 h1:fVRFRnXvU+x6C4IlHZewvJOVHoOv1TUuQyoRsYnB4bI= -google.golang.org/grpc v1.56.2/go.mod h1:I9bI3vqKfayGqPUAwGdOSu7kt6oIJLixfffKrpXqQ9s= -google.golang.org/protobuf v1.26.0-rc.1/go.mod h1:jlhhOSvTdKEhbULTjvd4ARK9grFBp09yW+WbY/TyQbw= -google.golang.org/protobuf v1.26.0/go.mod h1:9q0QmTI4eRPtz6boOQmLYwt+qCgq0jsYwAQnmE0givc= -google.golang.org/protobuf v1.31.0 h1:g0LDEJHgrBl9N9r17Ru3sqWhkIx2NB67okBHPwC7hs8= -google.golang.org/protobuf v1.31.0/go.mod h1:HV8QOd/L58Z+nl8r43ehVNZIU/HEI6OcFqwMG9pJV4I= -gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= -gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= -gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= -gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= -gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= +github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= +github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI= +github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY= +github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag= +github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE= +github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek= +github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps= +github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= +github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= +github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE= +github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg= +go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64= +go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y= +go.opentelemetry.io/otel v1.44.0 h1:JjwHmHpA4iZ3wBxluu2fbbE7j4kqlE8jXyAyPXH7HqU= +go.opentelemetry.io/otel v1.44.0/go.mod h1:BMgjTHL9WPRlRjL2oZCBTL4whCGtXch2H4BhOPIAyYc= +go.opentelemetry.io/otel/metric v1.44.0 h1:1w0gILTcHdr3YI+ixLyjemwrVnsMURbTZFrSYCdDdmc= +go.opentelemetry.io/otel/metric v1.44.0/go.mod h1:8O7hanEPBNgEMmybD3s2VBKcgWOCsA6tzHBPODAiquo= +go.opentelemetry.io/otel/sdk v1.44.0 h1:nHYwb9lK+fJPU/dnT6s7W7Z8itMWyqrnVfbheVYrZ58= +go.opentelemetry.io/otel/sdk v1.44.0/go.mod h1:Osuydd3Se74nqjAKxid74N5eC+jfEqfTegHRnq58oK0= +go.opentelemetry.io/otel/sdk/metric v1.44.0 h1:3LlKgI+VjbVsjNRFZJZAJ30WjXC5VkNRks6si09iEfI= +go.opentelemetry.io/otel/sdk/metric v1.44.0/go.mod h1:5B5pMARnXxKhltooO4xUuCBorl65a4EpnTalObqOigA= +go.opentelemetry.io/otel/trace v1.44.0 h1:jxF5CsGYCe74MCRx2X4g7WsY/VBKRqqpNvXlX/6gtIk= +go.opentelemetry.io/otel/trace v1.44.0/go.mod h1:oLl1jrMQAVo6v3GAggN+1VH9VIz9iUSvW53sW1Q8PIE= +go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw= +go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg= +golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To= +golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU= +golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= +golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8= +golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M= +gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4= +gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa h1:mZHHdPZl0dbGHCflZgAq/Q468DWVFcU2whhB2KAo8fk= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8= +google.golang.org/grpc v1.83.2 h1:EManeRomTObA0BU7I8vXgg/78uE5MJ9M8B39EX2WscU= +google.golang.org/grpc v1.83.2/go.mod h1:YPI1hK3kDked6iHvgX3tR0y+nX/qpMFKhPgFsokw1S8= +google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= +google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= diff --git a/hashring/hashring.go b/hashring/hashring.go index ede6b2c..0fef1bc 100644 --- a/hashring/hashring.go +++ b/hashring/hashring.go @@ -6,14 +6,14 @@ package hashring import ( + "cmp" "encoding/binary" "errors" "fmt" + "slices" "sort" "strings" "sync" - - "golang.org/x/exp/slices" ) var ( @@ -239,24 +239,12 @@ type virtualNode struct { members nodeRecord } -// compareUint64 should be replaced with the standard library's cmp.Compare once -// Go 1.21 is released. -func compareUint64(x, y uint64) int { - if x < y { - return -1 - } - if x > y { - return +1 - } - return 0 -} - func cmpVnode(a, b virtualNode) int { if a.hashvalue == b.hashvalue { if a.members.hashvalue == b.members.hashvalue { return strings.Compare(a.members.nodeKey, b.members.nodeKey) } - return compareUint64(a.members.hashvalue, b.members.hashvalue) + return cmp.Compare(a.members.hashvalue, b.members.hashvalue) } - return compareUint64(a.hashvalue, b.hashvalue) + return cmp.Compare(a.hashvalue, b.hashvalue) } diff --git a/hashring/hashring_test.go b/hashring/hashring_test.go index 88f0774..0c309f2 100644 --- a/hashring/hashring_test.go +++ b/hashring/hashring_test.go @@ -1,6 +1,7 @@ package hashring import ( + "encoding/binary" "errors" "fmt" "math" @@ -311,3 +312,174 @@ type member int func (m member) Key() string { return fmt.Sprintf("member-%d", m) } + +// tableHash is a HashFunc backed by a lookup table. +// It lets a test place members, virtual nodes and keys at exact positions on +// the ring, which a real hash function makes practically impossible. +// Hashing an input that is missing from the table fails the test. +type tableHash struct { + t *testing.T + hashes map[string]uint64 +} + +func (th tableHash) hash(b []byte) uint64 { + v, ok := th.hashes[string(b)] + if !ok { + th.t.Fatalf("no hash configured for input %q", b) + } + return v +} + +// vnodeInput mirrors the buffer Add hashes to place a virtual node: +// the member hash followed by the virtual node offset, both little-endian. +func vnodeInput(memberHash uint64, offset uint16) string { + buf := make([]byte, 10) + binary.LittleEndian.PutUint64(buf, memberHash) + binary.LittleEndian.PutUint16(buf[8:], offset) + return string(buf) +} + +func memberKeys(members []Member) []string { + keys := make([]string, 0, len(members)) + for _, m := range members { + keys = append(keys, m.Key()) + } + return keys +} + +// A key owned by the first virtual node at or after its hash. +// A key hashing exactly onto a virtual node belongs to that node, and a key +// hashing past the last virtual node wraps around to the first one. +func TestFindNWalksTheRingClockwiseFromTheKey(t *testing.T) { + th := tableHash{t, map[string]uint64{ + "a": 10, vnodeInput(10, 0): 100, + "b": 20, vnodeInput(20, 0): 200, + + "on-a": 100, + "between": 150, + "on-b": 200, + "past-end": 250, + }} + ring := MustNew(th.hash, 1) + require.NoError(t, ring.Add(testNode{nodeKeyAndValue: "a"})) + require.NoError(t, ring.Add(testNode{nodeKeyAndValue: "b"})) + + testCases := []struct { + key string + want []string + }{ + {"on-a", []string{"a", "b"}}, + {"between", []string{"b", "a"}}, + {"on-b", []string{"b", "a"}}, + {"past-end", []string{"a", "b"}}, + } + for _, tc := range testCases { + t.Run(tc.key, func(t *testing.T) { + one, err := ring.FindN([]byte(tc.key), 1) + require.NoError(t, err) + require.Equal(t, tc.want[:1], memberKeys(one)) + + two, err := ring.FindN([]byte(tc.key), 2) + require.NoError(t, err) + require.Equal(t, tc.want, memberKeys(two)) + }) + } +} + +// Two members whose keys hash to the same value produce identical virtual +// node hashes. +// They must still be told apart, so that removing one of them never removes +// the other's virtual nodes and the ring shape does not depend on the order +// in which they were added. +func TestMembersWithCollidingHashesStayDistinct(t *testing.T) { + const rf = 2 + th := tableHash{t, map[string]uint64{ + "a": 10, + "b": 10, + vnodeInput(10, 0): 100, + vnodeInput(10, 1): 300, + "k": 50, + }} + a, b := testNode{nodeKeyAndValue: "a"}, testNode{nodeKeyAndValue: "b"} + + for name, order := range map[string][]testNode{ + "a-then-b": {a, b}, + "b-then-a": {b, a}, + } { + t.Run(name, func(t *testing.T) { + ring := MustNew(th.hash, rf) + for _, m := range order { + require.NoError(t, ring.Add(m)) + } + require.Len(t, ring.virtualNodes, 2*rf) + + both, err := ring.FindN([]byte("k"), 2) + require.NoError(t, err) + require.ElementsMatch(t, []string{"a", "b"}, memberKeys(both)) + + // Insertion order must not change which member owns the key. + first, err := ring.FindN([]byte("k"), 1) + require.NoError(t, err) + require.Equal(t, []string{"a"}, memberKeys(first)) + + require.NoError(t, ring.Remove(a)) + require.Len(t, ring.virtualNodes, rf) + require.Equal(t, []string{"b"}, memberKeys(ring.Members())) + + left, err := ring.FindN([]byte("k"), 1) + require.NoError(t, err) + require.Equal(t, []string{"b"}, memberKeys(left)) + + require.ErrorIs(t, ring.Remove(a), ErrMemberNotFound) + }) + } +} + +func TestNewRejectsAReplicationFactorOfZero(t *testing.T) { + ring, err := New(xxhash.Sum64, 0) + require.ErrorIs(t, err, ErrInvalidReplicationFactor) + require.Nil(t, ring) + + ring, err = New(xxhash.Sum64, 1) + require.NoError(t, err) + require.NotNil(t, ring) +} + +func TestMustNewPanicsOnlyOnAnInvalidReplicationFactor(t *testing.T) { + require.PanicsWithError(t, ErrInvalidReplicationFactor.Error(), func() { + MustNew(xxhash.Sum64, 0) + }) + + var ring *Ring + require.NotPanics(t, func() { ring = MustNew(xxhash.Sum64, 1) }) + require.NotNil(t, ring) + require.Equal(t, uint16(1), ring.replicationFactor) +} + +// The ring is sorted and searched with cmpVnode, so it must be a strict total +// order with the same sign convention as cmp.Compare. +func TestCmpVnodeOrdersByHashThenMemberHashThenKey(t *testing.T) { + vn := func(hash, memberHash uint64, key string) virtualNode { + return virtualNode{hash, nodeRecord{hashvalue: memberHash, nodeKey: key}} + } + + testCases := []struct { + name string + a, b virtualNode + want int + }{ + {"lower vnode hash", vn(1, 9, "z"), vn(2, 1, "a"), -1}, + {"higher vnode hash", vn(2, 1, "a"), vn(1, 9, "z"), +1}, + {"same vnode hash, lower member hash", vn(5, 1, "z"), vn(5, 2, "a"), -1}, + {"same vnode hash, higher member hash", vn(5, 2, "a"), vn(5, 1, "z"), +1}, + {"same hashes, lower key", vn(5, 1, "a"), vn(5, 1, "b"), -1}, + {"same hashes, higher key", vn(5, 1, "b"), vn(5, 1, "a"), +1}, + {"identical", vn(5, 1, "a"), vn(5, 1, "a"), 0}, + } + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + require.Equal(t, tc.want, cmpVnode(tc.a, tc.b)) + require.Equal(t, -tc.want, cmpVnode(tc.b, tc.a)) + }) + } +}