diff --git a/etcd/embed.go b/etcd/embed.go index 59304c922..6afe9f547 100644 --- a/etcd/embed.go +++ b/etcd/embed.go @@ -78,10 +78,10 @@ const ( // We put all the things the node-watcher watches in /node so we // can use a single watcher for them. nodePrefix = "/node/" - heartbeatPrefix = "/node/heartbeat/" + heartbeatPrefix = nodePrefix + "heartbeat/" schemaPrefix = "/schema/" - resizePrefix = "/node/resize/" - metadataPrefix = "/node/metadata/" + resizePrefix = nodePrefix + "resize/" + metadataPrefix = nodePrefix + "metadata/" shardPrefix = "/shard/" lockPrefix = "/lock/" ) @@ -190,6 +190,12 @@ func (e *Etcd) Close() error { // retryClient attempts to do a thing, but also tries to handle the // specific case where the client fails because of a leader election, // in which case we need to restart the client and retry the thing. +// +// We have to let go of the lock while calling `fn` because some fn are +// long-lasting ones, like watchNodesOnce. So we grab a local copy of +// the client object, then call things on that object. This should error +// out sanely instead of panicing if we close the client while something +// is running on it. func (e *Etcd) retryClient(fn func(cli *clientv3.Client) error) (err error) { e.cliMu.Lock() cli := e.cli @@ -200,9 +206,14 @@ func (e *Etcd) retryClient(fn func(cli *clientv3.Client) error) (err error) { } // we can't do much with an error from closing e.cli at this point, so // we try again. + e.cliMu.Lock() + if cli != e.cli { + cli = e.cli + e.cliMu.Unlock() + return fn(cli) + } _ = cli.Close() cli = v3client.New(e.e.Server) - e.cliMu.Lock() e.cli = cli e.cliMu.Unlock() return fn(cli) @@ -312,7 +323,12 @@ func (e *Etcd) Start(ctx context.Context) (_ disco.InitialClusterState, err erro // heard from yet. for _, member := range members { peerID := member.ID.String() - e.knownNodes[peerID] = &nodeData{} + e.knownNodes[peerID] = &nodeData{ + topologyNode: &topology.Node{ + ID: peerID, + State: disco.NodeStateUnknown, + }, + } e.nodeStates[peerID] = disco.NodeStateUnknown } e.nodeStatesDirty = true @@ -324,7 +340,6 @@ func (e *Etcd) Start(ctx context.Context) (_ disco.InitialClusterState, err erro // watcher that watches for changes to events we care about. func (e *Etcd) startHeartbeatAndWatcher(ctx context.Context) error { key := heartbeatPrefix + e.e.Server.ID().String() - // e.logger.Printf("starting heartbeat: %q => %q", e.options.Name, e.e.Server.ID()) e.heartbeatLeasedKV = newLeasedKV(e, key, e.options.HeartbeatTTL) if err := e.heartbeatLeasedKV.Start(string(disco.NodeStateStarting)); err != nil { @@ -407,9 +422,6 @@ func (e *Etcd) ClusterState(ctx context.Context) (out disco.ClusterState, err er err = e.populateNodeStates(ctx) states := e.nodeStates e.nodeMu.Unlock() - // defer func() { - // e.logger.Printf("ClusterState %q: out %s, states %#v", e.options.Name, out, states) - // }() if err != nil { e.logger.Printf("ClusterState %q: getting node states: %v", e.options.Name, states) return disco.ClusterStateUnknown, err @@ -497,17 +509,17 @@ func (e *Etcd) Watch(ctx context.Context, peerID string, onUpdate func([]byte) e return nil } -// parseNodeKey reads "/node/heartbeat/23" and yields ("heartbeat", "23", nil). +// parseNodeKey reads heartbeatPrefix + "23" and yields (heartbeatPrefix, "23", nil). func parseNodeKey(key []byte) (prefix string, peerID string, err error) { - // we're looking for "/node/" - if !bytes.HasPrefix(key, []byte("/node/")) { + // we're looking for things starting with nodePrefix + if !bytes.HasPrefix(key, []byte(nodePrefix)) { return "", "", fmt.Errorf("not a node key: %q", key) } - vals := bytes.SplitN(key[6:], []byte("/"), 2) - if len(vals) != 2 { + peerIndex := bytes.LastIndex(key, []byte("/")) + if peerIndex < 6 { return "", "", fmt.Errorf("not a valid node key: %q", key) } - return string(vals[0]), string(vals[1]), nil + return string(key[:peerIndex+1]), string(key[peerIndex+1:]), nil } // deleteNodeData is like putNodeData, but handles deletes rather than cases @@ -521,20 +533,20 @@ func (e *Etcd) deleteNodeData(key []byte, revision int64) error { e.nodeRev = revision } switch prefix { - case "heartbeat": + case heartbeatPrefix: if e.knownNodes[peerID] == nil { e.knownNodes[peerID] = &nodeData{} } e.knownNodes[peerID].heartbeatState = "" e.nodeStatesDirty = true - case "metadata": + case metadataPrefix: if e.knownNodes[peerID] == nil { e.knownNodes[peerID] = &nodeData{} } e.knownNodes[peerID].metadata = nil e.knownNodes[peerID].topologyNode = &topology.Node{} e.nodeStatesDirty = true - case "resize": + case resizePrefix: if e.knownNodes[peerID] == nil { e.knownNodes[peerID] = &nodeData{} } @@ -550,9 +562,6 @@ func (e *Etcd) deleteNodeData(key []byte, revision int64) error { // given an incoming heartbeat, metadata, or resizing change. It requires // that you already hold the node mutex. func (e *Etcd) putNodeData(key []byte, value []byte, revision int64) (err error) { - // defer func() { - // e.logger.Printf("putNodeData %q: key %q, rev %d, dirty %t, err %v", e.options.Name, key, e.nodeRev, e.nodeStatesDirty, err) - // }() prefix, peerID, err := parseNodeKey(key) if err != nil { return err @@ -561,13 +570,13 @@ func (e *Etcd) putNodeData(key []byte, value []byte, revision int64) (err error) e.nodeRev = revision } switch prefix { - case "heartbeat": + case heartbeatPrefix: if e.knownNodes[peerID] == nil { e.knownNodes[peerID] = &nodeData{} } e.knownNodes[peerID].heartbeatState = string(value) e.nodeStatesDirty = true - case "metadata": + case metadataPrefix: if e.knownNodes[peerID] == nil { e.knownNodes[peerID] = &nodeData{} } @@ -577,12 +586,11 @@ func (e *Etcd) putNodeData(key []byte, value []byte, revision int64) (err error) if err != nil { return fmt.Errorf("json unmarshal of node metadata: %v\n", err) } - // e.logger.Printf("unmarshalling topologyNode for %q", peerID) e.knownNodes[peerID].topologyNode = &newNode // This saves us one remake of the node later, probably. e.knownNodes[peerID].topologyNode.State = e.knownNodes[peerID].computedState() e.nodeStatesDirty = true - case "resize": + case resizePrefix: if e.knownNodes[peerID] == nil { e.knownNodes[peerID] = &nodeData{} } @@ -601,20 +609,6 @@ func (e *Etcd) populateNodeStates(ctx context.Context) error { if !e.nodeStatesDirty { return nil } - updated := false - for peerID, data := range e.knownNodes { - newState := data.computedState() - if data.topologyNode == nil || newState != data.topologyNode.State || newState != e.nodeStates[peerID] { - updated = true - break - } - } - // no changes - if !updated { - e.nodeStatesDirty = false - return nil - } - // make a new node state map, because something changed. e.nodeStates = make(map[string]disco.NodeState, len(e.knownNodes)) e.sortedNodes = make([]*topology.Node, 0, len(e.knownNodes)) for peerID, data := range e.knownNodes { @@ -631,8 +625,6 @@ func (e *Etcd) populateNodeStates(ctx context.Context) error { data.topologyNode = &newNode } e.sortedNodes = append(e.sortedNodes, data.topologyNode) - } else { - e.logger.Printf("no topologyNode for peer %q, metadata %q", peerID, data.metadata) } } // sort list by ID. list now contains sorted nodes which have their @@ -650,11 +642,8 @@ func (e *Etcd) watchNodesOnce(ctx context.Context, cli *clientv3.Client) (err er // currently seen, we don't want one equal to it. minRev := e.nodeRev + 1 e.nodeMu.Unlock() - // e.logger.Printf("watchNodesOnce %q: minRev %d", e.options.Name, minRev) for resp := range cli.Watch(ctx, nodePrefix, clientv3.WithPrefix(), clientv3.WithRev(minRev)) { - // e.logger.Printf("watchp %q: resp rev %d with err %v, %d events\n", e.options.Name, resp.Header.Revision, resp.Err(), len(resp.Events)) if err := resp.Err(); err != nil { - e.logger.Printf("node watch %q: resp err %v", e.options.Name, err) return err } // lock the node mutex for this whole process of updating so @@ -664,13 +653,11 @@ func (e *Etcd) watchNodesOnce(ctx context.Context, cli *clientv3.Client) (err er for _, ev := range resp.Events { switch ev.Type { case mvccpb.PUT: - // e.logger.Printf("watchp %q: PUT rev %d value %q=%q\n", e.options.Name, ev.Kv.ModRevision, ev.Kv.Key, ev.Kv.Value) err := e.putNodeData(ev.Kv.Key, ev.Kv.Value, ev.Kv.ModRevision) if err != nil { e.logger.Printf("put event: %v", err) } case mvccpb.DELETE: - // e.logger.Printf("watchp %q: DELETE %q\n", e.options.Name, ev.Kv.Key) err := e.deleteNodeData(ev.Kv.Key, ev.Kv.ModRevision) if err != nil { e.logger.Printf("delete event: %v", err) @@ -690,7 +677,6 @@ func (e *Etcd) watchNodesOnce(ctx context.Context, cli *clientv3.Client) (err er func (e *Etcd) WatchNodes() { ctx, cancel := context.WithCancel(context.Background()) e.watchCancel = cancel - // e.logger.Printf("WatchNodes %q: starting with rev %d", e.options.Name, e.nodeRev) watchInContext := func(cli *clientv3.Client) error { return e.watchNodesOnce(ctx, cli) } @@ -1275,10 +1261,8 @@ func (e *Etcd) Nodes() []*topology.Node { defer e.nodeMu.Unlock() err := e.populateNodeStates(context.TODO()) if err != nil { - // e.logger.Printf("requesting node list: %v", err) return nil } - // e.logger.Printf("Nodes(): returning %d nodes", len(e.sortedNodes)) return e.sortedNodes } diff --git a/executor_test.go b/executor_test.go index 425a028a7..330ef1fd6 100644 --- a/executor_test.go +++ b/executor_test.go @@ -3764,7 +3764,7 @@ func TestReopenCluster(t *testing.T) { if err := node0.Reopen(); err != nil { t.Fatal(err) } - if err := node0.AwaitState(disco.ClusterStateNormal, 10*time.Second); err != nil { + if err := c.AwaitState(disco.ClusterStateNormal, 10*time.Second); err != nil { t.Fatalf("restarting cluster: %v", err) } @@ -3828,7 +3828,7 @@ func TestExecutor_Execute_Existence(t *testing.T) { t.Fatal(err) } - if err := node0.AwaitState(disco.ClusterStateNormal, 10*time.Second); err != nil { + if err := c.AwaitState(disco.ClusterStateNormal, 10*time.Second); err != nil { t.Fatalf("restarting cluster: %v", err) } diff --git a/server/cluster_test.go b/server/cluster_test.go index 97b592abf..16ae4f7a5 100644 --- a/server/cluster_test.go +++ b/server/cluster_test.go @@ -159,13 +159,8 @@ func TestClusterResize_AddNode(t *testing.T) { clus := test.MustRunCluster(t, 3) defer clus.Close() - state0, err0 := clus.GetNode(0).API.State() - state1, err1 := clus.GetNode(1).API.State() - if err0 != nil || !test.CheckClusterState(clus.GetNode(0), disco.ClusterStateNormal, 1000) { - t.Fatalf("unexpected node0 cluster state: %s, error: %v", state0, err0) - } else if err1 != nil || !test.CheckClusterState(clus.GetNode(1), disco.ClusterStateNormal, 1000) { - t.Fatalf("unexpected node1 cluster state: %s, error: %v", state1, err1) - } + clus.GetNode(0).AssertState(t, disco.ClusterStateNormal, 1*time.Second) + clus.GetNode(1).AssertState(t, disco.ClusterStateNormal, 1*time.Second) }) t.Run("WithIndex", func(t *testing.T) { // Configure node0 @@ -207,13 +202,8 @@ func TestClusterResize_AddNode(t *testing.T) { defer m1.Close() - state0, err0 := m0.API.State() - state1, err1 := m1.API.State() - if err0 != nil || !test.CheckClusterState(m0, disco.ClusterStateNormal, 1000) { - t.Fatalf("unexpected node0 cluster state: %s, error: %v", state0, err0) - } else if err1 != nil || !test.CheckClusterState(m1, disco.ClusterStateNormal, 1000) { - t.Fatalf("unexpected node1 cluster state: %s, error; %v", state1, err1) - } + m0.AssertState(t, disco.ClusterStateNormal, 1*time.Second) + m1.AssertState(t, disco.ClusterStateNormal, 1*time.Second) }) t.Run("ContinuousShards", func(t *testing.T) { @@ -253,13 +243,8 @@ func TestClusterResize_AddNode(t *testing.T) { m1 := c.GetNode(1) defer m1.Close() - state0, err0 := m0.API.State() - state1, err1 := m1.API.State() - if err0 != nil || !test.CheckClusterState(m0, disco.ClusterStateNormal, 1000) { - t.Fatalf("unexpected node0 cluster state: %s, error: %v", state0, err0) - } else if err1 != nil || !test.CheckClusterState(m1, disco.ClusterStateNormal, 1000) { - t.Fatalf("unexpected node1 cluster state: %s, error: %v", state1, err1) - } + m0.AssertState(t, disco.ClusterStateNormal, 1*time.Second) + m1.AssertState(t, disco.ClusterStateNormal, 1*time.Second) // Verify the data exists on both nodes. m0.QueryExpect(t, "i", "", `Row(f=1)`, exp) @@ -301,13 +286,8 @@ func TestClusterResize_AddNode(t *testing.T) { m1 := c.GetNode(1) defer m1.Close() - state0, err0 := m0.API.State() - state1, err1 := m1.API.State() - if err0 != nil || !test.CheckClusterState(m0, disco.ClusterStateNormal, 1000) { - t.Fatalf("unexpected node0 cluster state: %s, error: %v", state0, err0) - } else if err1 != nil || !test.CheckClusterState(m1, disco.ClusterStateNormal, 1000) { - t.Fatalf("unexpected node1 cluster state: %s, error: %v", state1, err1) - } + m0.AssertState(t, disco.ClusterStateNormal, 1*time.Second) + m1.AssertState(t, disco.ClusterStateNormal, 1*time.Second) // Verify the data exists on both nodes. m0.QueryExpect(t, "i", "", `Row(f=1)`, exp) @@ -353,13 +333,8 @@ func TestClusterResize_AddNode(t *testing.T) { m1 := c.GetNode(1) defer m1.Close() - state0, err0 := m0.API.State() - state1, err1 := m1.API.State() - if err0 != nil || !test.CheckClusterState(m0, disco.ClusterStateNormal, 1000) { - t.Fatalf("unexpected node0 cluster state: %s, error: %v", state0, err0) - } else if err1 != nil || !test.CheckClusterState(m1, disco.ClusterStateNormal, 1000) { - t.Fatalf("unexpected node1 cluster state: %s, error: %v", state1, err1) - } + m0.AssertState(t, disco.ClusterStateNormal, 1*time.Second) + m1.AssertState(t, disco.ClusterStateNormal, 1*time.Second) // Verify the data exists on both nodes. m0.QueryExpect(t, "i", "", `Row(f=1)`, exp) @@ -399,13 +374,8 @@ func TestClusterResize_AddNodeConcurrentIndex(t *testing.T) { m1 := c.GetNode(1) defer m1.Close() - state0, err0 := m0.API.State() - state1, err1 := m1.API.State() - if err0 != nil || !test.CheckClusterState(m0, disco.ClusterStateNormal, 1000) { - t.Fatalf("unexpected node0 cluster state: %s, error: %v", state0, err0) - } else if err1 != nil || !test.CheckClusterState(m1, disco.ClusterStateNormal, 1000) { - t.Fatalf("unexpected node1 cluster state: %s, error: %v", state1, err1) - } + m0.AssertState(t, disco.ClusterStateNormal, 1*time.Second) + m1.AssertState(t, disco.ClusterStateNormal, 1*time.Second) if err := <-errc; err != nil { t.Fatalf("error from index creation: %v", err) @@ -450,13 +420,8 @@ func TestClusterResize_AddNodeConcurrentIndex(t *testing.T) { m1 := c.GetNode(1) defer m1.Close() - state0, err0 := m0.API.State() - state1, err1 := m1.API.State() - if err0 != nil || !test.CheckClusterState(m0, disco.ClusterStateNormal, 1000) { - t.Fatalf("unexpected node0 cluster state: %s, error: %v", state0, err0) - } else if err1 != nil || !test.CheckClusterState(m1, disco.ClusterStateNormal, 1000) { - t.Fatalf("unexpected node1 cluster state: %s, error: %v", state1, err1) - } + m0.AssertState(t, disco.ClusterStateNormal, 1*time.Second) + m1.AssertState(t, disco.ClusterStateNormal, 1*time.Second) // Verify the data exists on both nodes. m0.QueryExpect(t, "i", "", `Row(f=1)`, exp) @@ -501,13 +466,8 @@ func TestClusterResize_AddNodeConcurrentIndex(t *testing.T) { m1 := c.GetNode(1) defer m1.Close() - state0, err0 := m0.API.State() - state1, err1 := m1.API.State() - if err0 != nil || !test.CheckClusterState(m0, disco.ClusterStateNormal, 1000) { - t.Fatalf("unexpected node0 cluster state: %s, error: %v", state0, err0) - } else if err1 != nil || !test.CheckClusterState(m1, disco.ClusterStateNormal, 1000) { - t.Fatalf("unexpected node1 cluster state: %s, error: %v", state1, err1) - } + m0.AssertState(t, disco.ClusterStateNormal, 1*time.Second) + m1.AssertState(t, disco.ClusterStateNormal, 1*time.Second) // Verify the data exists on both nodes. m0.QueryExpect(t, "i", "", `Row(f=1)`, exp) @@ -550,13 +510,9 @@ func TestClusterResize_AddNodeConcurrentIndex(t *testing.T) { m1 := c.GetNode(1) defer m1.Close() - state0, err0 := m0.API.State() - state1, err1 := m1.API.State() - if err0 != nil || !test.CheckClusterState(m0, disco.ClusterStateNormal, 1000) { - t.Fatalf("unexpected node0 cluster state: %s, error: %v", state0, err0) - } else if err1 != nil || !test.CheckClusterState(m1, disco.ClusterStateNormal, 1000) { - t.Fatalf("unexpected node1 cluster state: %s, error: %v", state1, err1) - } + m0.AssertState(t, disco.ClusterStateNormal, 1*time.Second) + m1.AssertState(t, disco.ClusterStateNormal, 1*time.Second) + m0.QueryExpect(t, "i", "", `Row(f=1)`, exp) m1.QueryExpect(t, "i", "", `Row(f=1)`, exp) }) diff --git a/server/server_test.go b/server/server_test.go index d2381f8de..771735f8b 100644 --- a/server/server_test.go +++ b/server/server_test.go @@ -371,7 +371,7 @@ func TestConcurrentFieldCreation(t *testing.T) { defer cluster.Close() node0 := cluster.GetNode(0) - err := node0.AwaitState(disco.ClusterStateNormal, 100*time.Millisecond) + err := cluster.AwaitState(disco.ClusterStateNormal, 100*time.Millisecond) if err != nil { t.Fatalf("starting cluster: %v", err) } @@ -643,7 +643,7 @@ func TestClusteringNodesReplica1(t *testing.T) { cluster := test.MustRunCluster(t, 3) defer cluster.Close() - if err := cluster.GetNode(0).AwaitState(disco.ClusterStateNormal, 100*time.Millisecond); err != nil { + if err := cluster.AwaitState(disco.ClusterStateNormal, 100*time.Millisecond); err != nil { t.Fatalf("starting cluster: %v", err) } @@ -651,7 +651,7 @@ func TestClusteringNodesReplica1(t *testing.T) { t.Fatalf("closing third node: %v", err) } - if err := cluster.GetPrimary().AwaitState(disco.ClusterStateDown, 30*time.Second); err != nil { + if err := cluster.AwaitPrimaryState(disco.ClusterStateDown, 30*time.Second); err != nil { t.Fatalf("starting cluster: %v", err) } @@ -681,7 +681,7 @@ func TestClusteringNodesReplica2(t *testing.T) { t.Fatalf("closing third node: %v", err) } - err = coord.AwaitState(disco.ClusterStateDegraded, 30*time.Second) + err = cluster.AwaitPrimaryState(disco.ClusterStateDegraded, 30*time.Second) if err != nil { t.Fatalf("after closing first server: %v", err) } @@ -699,7 +699,7 @@ func TestClusteringNodesReplica2(t *testing.T) { t.Fatalf("closing 2nd node: %v", err) } - err = coord.AwaitState(disco.ClusterStateDown, 30*time.Second) + err = cluster.AwaitPrimaryState(disco.ClusterStateDown, 30*time.Second) if err != nil { t.Fatalf("after closing second server: %v", err) } @@ -730,7 +730,7 @@ func TestRemoveNodeAfterItDies(t *testing.T) { coord, others := cluster.GetPrimary(), cluster.GetNonPrimaries() - err = coord.AwaitState(disco.ClusterStateNormal, 100*time.Millisecond) + err = cluster.AwaitState(disco.ClusterStateNormal, 100*time.Millisecond) if err != nil { t.Fatalf("starting cluster: %v", err) } @@ -741,16 +741,16 @@ func TestRemoveNodeAfterItDies(t *testing.T) { t.Fatalf("closing third node: %v", err) } - err = coord.AwaitState(disco.ClusterStateDegraded, 30*time.Second) + err = cluster.AwaitPrimaryState(disco.ClusterStateDegraded, 30*time.Second) if err != nil { - t.Fatalf("starting cluster: %v", err) + t.Fatalf("degrading cluster: %v", err) } if _, err := coord.API.RemoveNode(disabled.API.Node().ID); err != nil { t.Fatalf("removing failed node: %v", err) } - err = coord.AwaitState(disco.ClusterStateNormal, 30*time.Second) + err = cluster.AwaitPrimaryState(disco.ClusterStateNormal, 30*time.Second) if err != nil { t.Fatalf("removing disabled node: %v", err) } @@ -774,7 +774,7 @@ func TestRemoveConcurrentIndexCreation(t *testing.T) { defer cluster.Close() node0 := cluster.GetNode(0) - err = node0.AwaitState(disco.ClusterStateNormal, 100*time.Millisecond) + err = cluster.AwaitState(disco.ClusterStateNormal, 100*time.Millisecond) if err != nil { t.Fatalf("starting cluster: %v", err) } @@ -789,7 +789,7 @@ func TestRemoveConcurrentIndexCreation(t *testing.T) { t.Fatalf("removing node: %v", err) } - err = cluster.GetPrimary().AwaitState(disco.ClusterStateNormal, 100*time.Millisecond) + err = cluster.AwaitPrimaryState(disco.ClusterStateNormal, 100*time.Millisecond) if err != nil { t.Fatalf("starting cluster: %v", err) } diff --git a/test/cluster.go b/test/cluster.go index ec98f3e0f..22489bc9a 100644 --- a/test/cluster.go +++ b/test/cluster.go @@ -415,7 +415,7 @@ func (c *Cluster) Start() error { if err != nil { return errors.Wrap(err, "starting cluster") } - return c.GetNode(0).AwaitState(disco.ClusterStateNormal, 30*time.Second) + return c.AwaitState(disco.ClusterStateNormal, 30*time.Second) } // Close stops a Cluster @@ -447,6 +447,67 @@ func (c *Cluster) CloseAndRemove(n int) error { return err } +// AwaitPrimaryState waits for the cluster primary to reach a specified cluster state. +// When this happens, we know etcd reached a combination of node states that +// would imply this cluster state, but some nodes may not have caught up yet; +// we just test that the coordinator thought the cluster was in the given state. +func (c *Cluster) AwaitPrimaryState(expectedState disco.ClusterState, timeout time.Duration) error { + if len(c.Nodes) < 1 { + return errors.New("can't await coordinator state on an empty cluster") + } + primary := c.GetPrimary() + if primary == nil { + startTime := time.Now() + var elapsed time.Duration + for elapsed = 0; elapsed <= timeout; elapsed = time.Since(startTime) { + time.Sleep(50 * time.Millisecond) + primary = c.GetPrimary() + if primary != nil { + break + } + } + if primary == nil { + return errors.New("timed out waiting for cluster to have valid topology") + } + // we used up some of our timeout waiting for this + c.tb.Logf("had to wait %v for cluster topology", elapsed) + timeout -= elapsed + } + onlyCoordinator := &Cluster{Nodes: []*Command{primary}} + return onlyCoordinator.AwaitState(expectedState, timeout) +} + +// ExceptionalState returns an error if any node in the cluster is not +// in the expected state. +func (c *Cluster) ExceptionalState(expectedState disco.ClusterState) error { + for _, node := range c.Nodes { + state, err := node.API.State() + if err != nil || state != expectedState { + return fmt.Errorf("node %q: state %s: err %v", node.ID(), state, err) + } + } + return nil +} + +// AwaitState waits for the whole cluster to reach a specified state. +func (c *Cluster) AwaitState(expectedState disco.ClusterState, timeout time.Duration) (err error) { + if len(c.Nodes) < 1 { + return errors.New("can't await state of an empty cluster") + } + startTime := time.Now() + var elapsed time.Duration + for elapsed = 0; elapsed <= timeout; elapsed = time.Since(startTime) { + // Counterintuitive: We're returning if the err *is* nil, + // meaning we've reached the expected state. + if err = c.ExceptionalState(expectedState); err == nil { + return err + } + time.Sleep(50 * time.Millisecond) + } + return fmt.Errorf("waited %v for cluster to reach state %q: %v", + elapsed, expectedState, err) +} + // MustNewCluster creates a new cluster. If opts contains only one // slice of command options, those options are used with every node. // If it is empty, default options are used. Otherwise, it must contain size @@ -466,23 +527,6 @@ func MustNewCluster(tb testing.TB, size int, opts ...[]server.CommandOption) *Cl return c } -// CheckClusterState polls a given cluster for its state until it -// receives a matching state. It polls up to n times before returning. -func CheckClusterState(m *Command, state disco.ClusterState, n int) bool { - for i := 0; i < n; i++ { - - apiState, err := m.API.State() - if err != nil { - return false - } - if apiState == state { - return true - } - time.Sleep(100 * time.Millisecond) - } - return false -} - // newCluster creates a new cluster func newCluster(tb testing.TB, size int, opts ...[]server.CommandOption) (*Cluster, error) { if size == 0 { diff --git a/test/pilosa.go b/test/pilosa.go index e9c19f552..66b17ea4e 100644 --- a/test/pilosa.go +++ b/test/pilosa.go @@ -384,12 +384,29 @@ func (m *Command) AwaitState(expectedState disco.ClusterState, timeout time.Dura if err = m.exceptionalState(expectedState); err == nil { return err } - time.Sleep(1 * time.Millisecond) + time.Sleep(50 * time.Millisecond) } return fmt.Errorf("waited %v for command to reach state %q: %v", elapsed, expectedState, err) } +// AssertState waits for the whole cluster to reach a specified state, or +// fails the calling test if it can't. +func (m *Command) AssertState(t testing.TB, expectedState disco.ClusterState, timeout time.Duration) { + startTime := time.Now() + var elapsed time.Duration + var err error + for elapsed = 0; elapsed <= timeout; elapsed = time.Since(startTime) { + // Counterintuitive: We're returning if the err *is* nil, + // meaning we've reached the expected state. + if err = m.exceptionalState(expectedState); err == nil { + return + } + time.Sleep(50 * time.Millisecond) + } + t.Fatalf("waited %v for command to reach state %q: %v", elapsed, expectedState, err) +} + // exceptionalState returns an error if the node is not in the expected state. func (m *Command) exceptionalState(expectedState disco.ClusterState) error { state, err := m.API.State()