diff --git a/pkg/keyspace/tso_keyspace_group.go b/pkg/keyspace/tso_keyspace_group.go index 08bec239d85..90e97c36cd7 100644 --- a/pkg/keyspace/tso_keyspace_group.go +++ b/pkg/keyspace/tso_keyspace_group.go @@ -58,6 +58,7 @@ const ( defaultKeyspaceCountSplitThreshold = 40000 // autoSplitKeyspaceGroupPatrolInterval is the interval for patrolling keyspace group size for auto-split. autoSplitKeyspaceGroupPatrolInterval = 15 * time.Minute + keyspaceGroupRevisionValue = "1" ) const ( @@ -90,6 +91,58 @@ type GroupManager struct { tsoNodesWatcher *etcdutil.LoopWatcher } +type keyspaceGroupRevisionStorage struct { + endpoint.KeyspaceGroupStorage +} + +type keyspaceGroupRevisionTxn struct { + kv.Txn + changed bool +} + +// Save records keyspace group writes so the transaction can advance the revision marker. +func (txn *keyspaceGroupRevisionTxn) Save(key, value string) error { + if err := txn.Txn.Save(key, value); err != nil { + return err + } + if strings.HasPrefix(key, keypath.KeyspaceGroupIDPrefix()) { + txn.changed = true + } + return nil +} + +// Remove records keyspace group deletions so the transaction can advance the revision marker. +func (txn *keyspaceGroupRevisionTxn) Remove(key string) error { + if err := txn.Txn.Remove(key); err != nil { + return err + } + if strings.HasPrefix(key, keypath.KeyspaceGroupIDPrefix()) { + txn.changed = true + } + return nil +} + +// RunInTxn atomically advances the revision marker when the transaction changes a keyspace group. +func (s *keyspaceGroupRevisionStorage) RunInTxn(ctx context.Context, f func(txn kv.Txn) error) error { + return s.KeyspaceGroupStorage.RunInTxn(ctx, func(txn kv.Txn) error { + revisionTxn := &keyspaceGroupRevisionTxn{Txn: txn} + if err := f(revisionTxn); err != nil { + return err + } + if !revisionTxn.changed { + return nil + } + return txn.Save(keypath.KeyspaceGroupRevisionPath(), keyspaceGroupRevisionValue) + }) +} + +func withKeyspaceGroupRevision(store endpoint.KeyspaceGroupStorage) endpoint.KeyspaceGroupStorage { + if _, ok := store.(*keyspaceGroupRevisionStorage); ok { + return store + } + return &keyspaceGroupRevisionStorage{KeyspaceGroupStorage: store} +} + // NewKeyspaceGroupManager creates a Manager of keyspace group related data. func NewKeyspaceGroupManager( ctx context.Context, @@ -104,7 +157,7 @@ func NewKeyspaceGroupManager( m := &GroupManager{ ctx: ctx, cancel: cancel, - store: store, + store: withKeyspaceGroupRevision(store), groups: groups, client: client, nodesBalancer: balancer.GenByPolicy[string](defaultBalancerPolicy), @@ -141,6 +194,13 @@ func (m *GroupManager) Bootstrap(ctx context.Context) error { if err != nil && err != errs.ErrKeyspaceGroupExists { return err } + // Persist a marker on every bootstrap so upgraded clusters retain a + // revision that is at least as new as all existing keyspace group state. + if err := m.store.RunInTxn(ctx, func(txn kv.Txn) error { + return txn.Save(keypath.KeyspaceGroupRevisionPath(), keyspaceGroupRevisionValue) + }); err != nil { + return err + } // Load all the keyspace groups from the storage and add to the respective userKind groups. groups, err := m.store.LoadKeyspaceGroups(constant.DefaultKeyspaceGroupID, 0) diff --git a/pkg/keyspace/tso_keyspace_group_test.go b/pkg/keyspace/tso_keyspace_group_test.go index 489dbbf7efd..6a950001dc7 100644 --- a/pkg/keyspace/tso_keyspace_group_test.go +++ b/pkg/keyspace/tso_keyspace_group_test.go @@ -33,6 +33,7 @@ import ( "github.com/tikv/pd/pkg/storage/endpoint" "github.com/tikv/pd/pkg/storage/kv" "github.com/tikv/pd/pkg/utils/etcdutil" + "github.com/tikv/pd/pkg/utils/keypath" "github.com/tikv/pd/pkg/versioninfo/kerneltype" ) @@ -131,6 +132,37 @@ func (suite *keyspaceGroupTestSuite) TestKeyspaceGroupOperations() { re.Error(err) } +func (suite *keyspaceGroupTestSuite) TestKeyspaceGroupChangesUpdateRevisionMarker() { + re := suite.Require() + + removeMarker := func() { + re.NoError(suite.kgm.store.RunInTxn(suite.ctx, func(txn kv.Txn) error { + return txn.Remove(keypath.KeyspaceGroupRevisionPath()) + })) + } + checkMarker := func() { + var marker string + re.NoError(suite.kgm.store.RunInTxn(suite.ctx, func(txn kv.Txn) error { + var err error + marker, err = txn.Load(keypath.KeyspaceGroupRevisionPath()) + return err + })) + re.NotEmpty(marker) + } + + removeMarker() + re.NoError(suite.kgm.CreateKeyspaceGroups([]*endpoint.KeyspaceGroup{{ + ID: 1, + UserKind: endpoint.Standard.String(), + }})) + checkMarker() + + removeMarker() + _, err := suite.kgm.DeleteKeyspaceGroupByID(1) + re.NoError(err) + checkMarker() +} + func (suite *keyspaceGroupTestSuite) TestKeyspaceAssignment() { re := suite.Require() diff --git a/pkg/tso/keyspace_group_manager.go b/pkg/tso/keyspace_group_manager.go index e3564a84616..8f1afcf1b40 100644 --- a/pkg/tso/keyspace_group_manager.go +++ b/pkg/tso/keyspace_group_manager.go @@ -534,9 +534,25 @@ func (kgm *KeyspaceGroupManager) InitializeTSOServerWatchLoop() error { // membership/distribution metadata. // Key: /pd/{cluster_id}/tso/keyspace_groups/membership/{group} // Value: endpoint.KeyspaceGroup +// Revision marker: /pd/{cluster_id}/tso/keyspace_groups/revision func (kgm *KeyspaceGroupManager) InitializeGroupWatchLoop() error { defaultKGConfigured := false + maxLoadedModRevision := uint64(0) + preEventsFn := func([]*clientv3.Event) error { + maxLoadedModRevision = 0 + return nil + } putFn := func(kv *mvccpb.KeyValue) error { + if string(kv.Key) == keypath.KeyspaceGroupRevisionPath() { + failpoint.Inject("SkipKeyspaceWatch", func(val failpoint.Value) { + addr, ok := val.(string) + if ok && addr == kgm.electionNamePrefix { + failpoint.Return(nil) + } + }) + maxLoadedModRevision = max(maxLoadedModRevision, uint64(kv.ModRevision)) + return nil + } group := &endpoint.KeyspaceGroup{} if err := json.Unmarshal(kv.Value, group); err != nil { return errs.ErrJSONUnmarshal.Wrap(err) @@ -547,6 +563,7 @@ func (kgm *KeyspaceGroupManager) InitializeGroupWatchLoop() error { failpoint.Return(nil) } }) + maxLoadedModRevision = max(maxLoadedModRevision, uint64(kv.ModRevision)) kgm.updateKeyspaceGroup(group) if group.ID == constant.DefaultKeyspaceGroupID { defaultKGConfigured = true @@ -554,6 +571,9 @@ func (kgm *KeyspaceGroupManager) InitializeGroupWatchLoop() error { return nil } deleteFn := func(kv *mvccpb.KeyValue) error { + if string(kv.Key) == keypath.KeyspaceGroupRevisionPath() { + return nil + } groupID, err := ExtractKeyspaceGroupIDFromPath(kgm.compiledKGMembershipIDRegexp, string(kv.Key)) if err != nil { return err @@ -581,6 +601,8 @@ func (kgm *KeyspaceGroupManager) InitializeGroupWatchLoop() error { zap.Uint64("new-mod-revision", uint64(last.Kv.ModRevision)), ) } + } else if maxLoadedModRevision > 0 { + kgm.SetModRevision(maxLoadedModRevision) } return nil } @@ -589,9 +611,8 @@ func (kgm *KeyspaceGroupManager) InitializeGroupWatchLoop() error { &kgm.wg, kgm.etcdClient, "keyspace-watcher", - // To keep the consistency with the previous code, we should trim the suffix `/`. - strings.TrimSuffix(keypath.KeyspaceGroupIDPrefix(), "/"), - func([]*clientv3.Event) error { return nil }, + keypath.KeyspaceGroupPrefix(), + preEventsFn, putFn, deleteFn, postEventsFn, diff --git a/pkg/tso/keyspace_group_manager_test.go b/pkg/tso/keyspace_group_manager_test.go index 12c4b217cad..195efd558ba 100644 --- a/pkg/tso/keyspace_group_manager_test.go +++ b/pkg/tso/keyspace_group_manager_test.go @@ -210,6 +210,140 @@ func (suite *keyspaceGroupManagerTestSuite) TestLoadKeyspaceGroupsAssignment() { suite.runTestLoadKeyspaceGroupsAssignment(re, maxCountInUse+1, 0, 10) } +func (suite *keyspaceGroupManagerTestSuite) TestLoadKeyspaceGroupsSetsModRevision() { + re := suite.Require() + + mgr := suite.newUniqueKeyspaceGroupManager(1) + re.NotNil(mgr) + defer mgr.Close() + + const ( + groupID = uint32(1) + keyspaceID = uint32(101) + ) + err := addKeyspaceGroupAssignment( + suite.ctx, + suite.etcdClient, + groupID, + []string{mgr.tsoServiceID.ServiceAddr}, + []int{mcs.DefaultKeyspaceGroupReplicaPriority}, + []uint32{keyspaceID}, + ) + re.NoError(err) + + resp, err := suite.etcdClient.Get(suite.ctx, keypath.KeyspaceGroupIDPath(groupID)) + re.NoError(err) + re.Len(resp.Kvs, 1) + groupRevision := uint64(resp.Kvs[0].ModRevision) + + const deletedGroupID = uint32(2) + err = addKeyspaceGroupAssignment( + suite.ctx, + suite.etcdClient, + deletedGroupID, + []string{mgr.tsoServiceID.ServiceAddr}, + []int{mcs.DefaultKeyspaceGroupReplicaPriority}, + []uint32{keyspaceID + 1}, + ) + re.NoError(err) + deleteResp, err := suite.etcdClient.Txn(suite.ctx).Then( + clientv3.OpDelete(keypath.KeyspaceGroupIDPath(deletedGroupID)), + clientv3.OpPut(keypath.KeyspaceGroupRevisionPath(), "1"), + ).Commit() + re.NoError(err) + deletedRevision := uint64(deleteResp.Header.Revision) + re.Greater(deletedRevision, groupRevision) + + err = mgr.Initialize() + re.NoError(err) + + _, kg, loadedGroupID, loadedRevision, err := mgr.FindGroupByKeyspaceID(keyspaceID) + re.NoError(err) + re.NotNil(kg) + re.Equal(groupID, loadedGroupID) + re.Equal(deletedRevision, loadedRevision) +} + +func (suite *keyspaceGroupManagerTestSuite) TestSnapshotRevisionRemainsComparableAcrossManagers() { + re := suite.Require() + keypath.SetClusterID(rand.Uint64()) + + cfg1 := suite.createConfig() + cfg2 := suite.createConfig() + mgr1 := suite.newKeyspaceGroupManager(1, cfg1) + mgr2 := suite.newKeyspaceGroupManager(1, cfg2) + defer mgr1.Close() + defer mgr2.Close() + + const ( + groupID = uint32(1) + keyspaceID = uint32(101) + ) + re.NoError(addKeyspaceGroupAssignment( + suite.ctx, + suite.etcdClient, + groupID, + []string{ + mgr1.tsoServiceID.ServiceAddr, + mgr2.tsoServiceID.ServiceAddr, + }, + []int{ + mcs.DefaultKeyspaceGroupReplicaPriority, + mcs.DefaultKeyspaceGroupReplicaPriority, + }, + []uint32{keyspaceID}, + )) + re.NoError(mgr1.Initialize()) + + _, _, _, oldRevision, err := mgr1.FindGroupByKeyspaceID(keyspaceID) + re.NoError(err) + re.NotZero(oldRevision) + + _, err = suite.etcdClient.Put(suite.ctx, "/unrelated/revision-gap", "1") + re.NoError(err) + + re.NoError(mgr2.Initialize()) + _, _, _, newRevision, err := mgr2.FindGroupByKeyspaceID(keyspaceID) + re.NoError(err) + re.Equal(oldRevision, newRevision) +} + +func (suite *keyspaceGroupManagerTestSuite) TestInitialSkipKeyspaceWatchDoesNotAdvanceRevision() { + re := suite.Require() + + mgr := suite.newUniqueKeyspaceGroupManager(1) + re.NotNil(mgr) + defer mgr.Close() + + point := fmt.Sprintf("return(\"%s\")", mgr.electionNamePrefix) + re.NoError(failpoint.Enable("github.com/tikv/pd/pkg/tso/SkipKeyspaceWatch", point)) + defer func() { + re.NoError(failpoint.Disable("github.com/tikv/pd/pkg/tso/SkipKeyspaceWatch")) + }() + + const ( + groupID = uint32(1) + keyspaceID = uint32(101) + ) + re.NoError(addKeyspaceGroupAssignment( + suite.ctx, + suite.etcdClient, + groupID, + []string{mgr.tsoServiceID.ServiceAddr}, + []int{mcs.DefaultKeyspaceGroupReplicaPriority}, + []uint32{keyspaceID}, + )) + re.NoError(mgr.Initialize()) + + mgr.RLock() + loadedGroup := mgr.kgs[groupID] + loadedRevision := mgr.modRevision + mgr.RUnlock() + + re.Nil(loadedGroup) + re.Zero(loadedRevision, "a skipped initial watch must not advertise an unapplied revision") +} + // TestLoadWithDifferentBatchSize tests the loading of the keyspace group assignment with the different batch size. func (suite *keyspaceGroupManagerTestSuite) TestLoadWithDifferentBatchSize() { re := suite.Require() diff --git a/pkg/utils/etcdutil/etcdutil.go b/pkg/utils/etcdutil/etcdutil.go index aa9f729aa7a..7b5e145f4a8 100644 --- a/pkg/utils/etcdutil/etcdutil.go +++ b/pkg/utils/etcdutil/etcdutil.go @@ -374,7 +374,6 @@ type LoopWatcher struct { postEventsFn func([]*clientv3.Event) error // preEventsFn is used to call before handling all events. preEventsFn func([]*clientv3.Event) error - // forceLoadMu is used to ensure two force loads have minimal interval. forceLoadMu syncutil.RWMutex // lastTimeForceLoad is used to record the last time force loading data from etcd. @@ -639,6 +638,7 @@ func (lw *LoopWatcher) watch(ctx context.Context, revision int64) (nextRevision func (lw *LoopWatcher) load(ctx context.Context) (nextRevision int64, err error) { startKey := lw.key limit := lw.loadBatchSize + snapshotRevision := int64(0) opts := lw.buildLoadingOpts(limit) if err := lw.preEventsFn([]*clientv3.Event{}); err != nil { @@ -677,10 +677,19 @@ func (lw *LoopWatcher) load(ctx context.Context) (nextRevision int64, err error) return 0, err } opts = lw.buildLoadingOpts(limit) + if snapshotRevision > 0 { + opts = append(opts, clientv3.WithRev(snapshotRevision)) + } continue } return 0, err } + if snapshotRevision == 0 { + snapshotRevision = resp.Header.Revision + // Keep all remaining pages on the same snapshot. Otherwise a write + // between pages could be skipped when the watch starts. + opts = append(opts, clientv3.WithRev(snapshotRevision)) + } for i, item := range resp.Kvs { if i == len(resp.Kvs)-1 && resp.More { // If there are more keys, we need to load the next batch. @@ -701,7 +710,7 @@ func (lw *LoopWatcher) load(ctx context.Context) (nextRevision int64, err error) } // Note: if there are no keys in etcd, the resp.More is false. It also means the load is finished. if !resp.More { - return resp.Header.Revision + 1, err + return snapshotRevision + 1, err } } } diff --git a/pkg/utils/etcdutil/etcdutil_test.go b/pkg/utils/etcdutil/etcdutil_test.go index 970c42325b5..631df07d59c 100644 --- a/pkg/utils/etcdutil/etcdutil_test.go +++ b/pkg/utils/etcdutil/etcdutil_test.go @@ -467,6 +467,48 @@ func (suite *loopWatcherTestSuite) TestLoadNoExistedKey() { re.Empty(cache) } +func (suite *loopWatcherTestSuite) TestLoadUsesSingleSnapshotRevision() { + re := suite.Require() + ctx, cancel := context.WithCancel(suite.ctx) + defer cancel() + + prefix := "TestLoadUsesSingleSnapshotRevision/" + _, err := suite.client.Txn(ctx).Then( + clientv3.OpPut(prefix+"a", ""), + clientv3.OpPut(prefix+"b", ""), + ).Commit() + re.NoError(err) + resp, err := suite.client.Get(ctx, prefix, clientv3.WithPrefix()) + re.NoError(err) + re.Len(resp.Kvs, 2) + expectedRevision := resp.Header.Revision + + var once sync.Once + var putErr error + watcher := NewLoopWatcher( + ctx, + &suite.wg, + suite.client, + "test", + prefix, + func([]*clientv3.Event) error { return nil }, + func(*mvccpb.KeyValue) error { + once.Do(func() { + _, putErr = suite.client.Put(ctx, prefix+"c", "") + }) + return putErr + }, + func(*mvccpb.KeyValue) error { return nil }, + func([]*clientv3.Event) error { return nil }, + true, /* withPrefix */ + ) + watcher.SetLoadBatchSize(1) + nextRevision, err := watcher.load(ctx) + re.NoError(err) + re.NoError(putErr) + re.Equal(expectedRevision+1, nextRevision) +} + func (suite *loopWatcherTestSuite) TestLoadWithLimitChange() { re := suite.Require() re.NoError(failpoint.Enable("github.com/tikv/pd/pkg/utils/etcdutil/meetEtcdError", `return()`)) diff --git a/pkg/utils/keypath/absolute_key_path.go b/pkg/utils/keypath/absolute_key_path.go index 28ca4e8dfaf..9cf23304d52 100644 --- a/pkg/utils/keypath/absolute_key_path.go +++ b/pkg/utils/keypath/absolute_key_path.go @@ -92,12 +92,14 @@ const ( minResolvedTSPathFormat = "/pd/%d/raft/min_resolved_ts" // "/pd/{cluster_id}/raft/min_resolved_ts" externalTimestampPathFormat = "/pd/%d/raft/external_timestamp" // "/pd/{cluster_id}/raft/external_timestamp" - keyspaceMetaPrefixFormat = "/pd/%d/keyspaces/meta/" // "/pd/{cluster_id}/keyspaces/meta/" - keyspaceMetaPathFormat = "/pd/%d/keyspaces/meta/%08d" // "/pd/{cluster_id}/keyspaces/meta/{keyspace_id}" - keyspaceIDPathFormat = "/pd/%d/keyspaces/id/%s" // "/pd/{cluster_id}/keyspaces/id/{keyspace_name}" - keyspaceGroupIDPrefixFormat = "/pd/%d/tso/keyspace_groups/membership/" // "/pd/{cluster_id}/tso/keyspace_groups/membership/" - keyspaceGroupIDPathFormat = "/pd/%d/tso/keyspace_groups/membership/%05d" // "/pd/{cluster_id}/tso/keyspace_groups/membership/{group_id}" - keyspaceGroupIDPattern = `tso/keyspace_groups/membership/(\d{5})$` + keyspaceMetaPrefixFormat = "/pd/%d/keyspaces/meta/" // "/pd/{cluster_id}/keyspaces/meta/" + keyspaceMetaPathFormat = "/pd/%d/keyspaces/meta/%08d" // "/pd/{cluster_id}/keyspaces/meta/{keyspace_id}" + keyspaceIDPathFormat = "/pd/%d/keyspaces/id/%s" // "/pd/{cluster_id}/keyspaces/id/{keyspace_name}" + keyspaceGroupPrefixFormat = "/pd/%d/tso/keyspace_groups/" // "/pd/{cluster_id}/tso/keyspace_groups/" + keyspaceGroupIDPrefixFormat = "/pd/%d/tso/keyspace_groups/membership/" // "/pd/{cluster_id}/tso/keyspace_groups/membership/" + keyspaceGroupIDPathFormat = "/pd/%d/tso/keyspace_groups/membership/%05d" // "/pd/{cluster_id}/tso/keyspace_groups/membership/{group_id}" + keyspaceGroupRevisionPathFormat = "/pd/%d/tso/keyspace_groups/revision" // "/pd/{cluster_id}/tso/keyspace_groups/revision" + keyspaceGroupIDPattern = `tso/keyspace_groups/membership/(\d{5})$` servicePathFormat = "/ms/%d/%s/registry/" // "/ms/{cluster_id}/{service_name}/registry/" registryPathFormat = "/ms/%d/%s/registry/%s" // "/ms/{cluster_id}/{service_name}/registry/{service_addr}" @@ -231,6 +233,11 @@ func KeyspaceIDPath(name string) string { return fmt.Sprintf(keyspaceIDPathFormat, ClusterID(), name) } +// KeyspaceGroupPrefix returns the prefix of keyspace group metadata. +func KeyspaceGroupPrefix() string { + return fmt.Sprintf(keyspaceGroupPrefixFormat, ClusterID()) +} + // KeyspaceGroupIDPrefix returns the prefix of keyspace group id. func KeyspaceGroupIDPrefix() string { return fmt.Sprintf(keyspaceGroupIDPrefixFormat, ClusterID()) @@ -241,6 +248,11 @@ func KeyspaceGroupIDPath(id uint32) string { return fmt.Sprintf(keyspaceGroupIDPathFormat, ClusterID(), id) } +// KeyspaceGroupRevisionPath returns the path of the durable keyspace group revision marker. +func KeyspaceGroupRevisionPath() string { + return fmt.Sprintf(keyspaceGroupRevisionPathFormat, ClusterID()) +} + // GetCompiledKeyspaceGroupIDRegexp returns the compiled regular expression for matching keyspace group id. func GetCompiledKeyspaceGroupIDRegexp() *regexp.Regexp { return regexp.MustCompile(keyspaceGroupIDPattern) diff --git a/server/server.go b/server/server.go index fe0a2a37d1c..2f0782d4829 100644 --- a/server/server.go +++ b/server/server.go @@ -671,28 +671,87 @@ func (s *Server) IsClosed() bool { // Run runs the pd server. func (s *Server) Run() error { + return s.RunWithContext(s.ctx) +} + +// RunWithContext runs the PD server with a context that controls startup and +// the server loops. The caller must keep the context alive while the server is +// running and call Close after canceling it. +func (s *Server) RunWithContext(ctx context.Context) (retErr error) { go systimemon.StartMonitor(s.ctx, time.Now, func() { log.Error("system time jumps backward", errs.ZapError(errs.ErrIncorrectSystemTime)) timeJumpBackCounter.Inc() }) - if err := s.startEtcd(s.ctx); err != nil { + if err := s.startEtcd(ctx); err != nil { return err } + defer func() { + if retErr != nil { + s.cleanupFailedStart() + } + }() + failpoint.Inject("failAfterStartEtcd", func() { + failpoint.Return(errors.New("injected error after etcd startup")) + }) - if err := s.startServer(s.ctx); err != nil { + if err := s.startServer(ctx); err != nil { return err } - s.cgMonitor.StartMonitor(s.ctx) + s.cgMonitor.StartMonitor(ctx) failpoint.Inject("delayStartServerLoop", func() { time.Sleep(2 * time.Second) }) - s.startServerLoop(s.ctx) + s.startServerLoop(ctx) return nil } +func (s *Server) cleanupFailedStart() { + if s.cluster != nil { + s.cluster.Stop() + } + if s.IsKeyspaceGroupEnabled() && s.keyspaceGroupManager != nil { + s.keyspaceGroupManager.Close() + } + if s.tsoAllocator != nil { + s.tsoAllocator.Close() + } + if s.meteringWriter != nil { + s.meteringWriter.Stop() + } + if s.client != nil { + if err := s.client.Close(); err != nil { + log.Warn("close etcd client meet error", errs.ZapError(errs.ErrCloseEtcdClient, err)) + } + } + if s.electionClient != nil { + if err := s.electionClient.Close(); err != nil { + log.Warn("close election client meet error", errs.ZapError(errs.ErrCloseEtcdClient, err)) + } + } + if s.httpClient != nil { + s.httpClient.CloseIdleConnections() + } + if s.member.Etcd() != nil { + s.member.Close() + } + if s.hbStreams != nil { + s.hbStreams.Close() + } + if s.storage != nil { + if err := s.storage.Close(); err != nil { + log.Warn("close storage meet error", errs.ZapError(err)) + } + } + if s.hotRegionStorage != nil { + if err := s.hotRegionStorage.Close(); err != nil { + log.Warn("close hot region storage meet error", errs.ZapError(err)) + } + } +} + // SetServiceAuditBackendForHTTP is used to register service audit config for HTTP. func (s *Server) SetServiceAuditBackendForHTTP(route *mux.Route, labels ...string) { if len(route.GetName()) == 0 { diff --git a/tests/cluster.go b/tests/cluster.go index f14cde38dd1..31abadc20fc 100644 --- a/tests/cluster.go +++ b/tests/cluster.go @@ -16,6 +16,7 @@ package tests import ( "context" + stdErrors "errors" "net/http" "os" "strings" @@ -77,6 +78,13 @@ var ( // defaultMaxRetryTimes is the default maximum retry times for starting servers. defaultMaxRetryTimes = 5 + // runServersCleanupGracePeriod allows the remaining servers to finish startup + // naturally after a peer fails. Canceling etcd before it becomes ready can + // leave its embedded listeners blocked in shutdown. + runServersCleanupGracePeriod = 10 * time.Second + // errRunServersCleanupTimeout prevents retry cleanup from destroying a data + // directory while its startup goroutine may still be using it. + errRunServersCleanupTimeout = errors.New("timed out waiting for starting servers to stop") ) type startServersRetryAction int @@ -99,6 +107,9 @@ func classifyInitialServersError(err error) startServersRetryAction { if err == nil { return startServersNoRetry } + if stdErrors.Is(err, errRunServersCleanupTimeout) { + return startServersNoRetry + } errMsg := err.Error() switch { case strings.Contains(errMsg, "address already in use") || strings.Contains(errMsg, "Etcd cluster ID mismatch"): @@ -116,6 +127,8 @@ type TestServer struct { server *server.Server grpcServer *server.GrpcServer state int32 + runCancel context.CancelFunc + runDone chan struct{} } var zapLogOnce sync.Once @@ -171,18 +184,56 @@ func NewTestServer(ctx context.Context, cfg *config.Config, services []string, h // Run starts to run a TestServer. func (s *TestServer) Run() error { + return s.runWithStartSignal(nil) +} + +func (s *TestServer) runWithStartSignal(started chan<- struct{}) error { s.Lock() - defer s.Unlock() if s.state != Initial && s.state != Stop { - return errors.Errorf("server(state%d) cannot run", s.state) + state := s.state + s.Unlock() + if started != nil { + close(started) + } + return errors.Errorf("server(state%d) cannot run", state) } - if err := s.server.Run(); err != nil { + prevState := s.state + runCtx, runCancel := context.WithCancel(s.server.Context()) + runDone := make(chan struct{}) + // Treat startup as running so retry cleanup can close a blocked server.Run. + s.state = Running + s.runCancel = runCancel + s.runDone = runDone + s.Unlock() + if started != nil { + close(started) + } + + err := s.server.RunWithContext(runCtx) + close(runDone) + if err != nil { + s.Lock() + if s.state == Running { + s.state = prevState + runCancel() + s.runCancel = nil + s.runDone = nil + } + s.Unlock() return err } - s.state = Running return nil } +func (s *TestServer) cancelRun() { + s.RLock() + runCancel := s.runCancel + s.RUnlock() + if runCancel != nil { + runCancel() + } +} + // Stop is used to stop a TestServer. func (s *TestServer) Stop() error { s.Lock() @@ -190,7 +241,15 @@ func (s *TestServer) Stop() error { if s.state != Running { return errors.Errorf("server(state%d) cannot stop", s.state) } + if s.runCancel != nil { + s.runCancel() + } + if s.runDone != nil { + <-s.runDone + } s.server.Close() + s.runCancel = nil + s.runDone = nil s.state = Stop return nil } @@ -200,7 +259,15 @@ func (s *TestServer) Destroy() error { s.Lock() defer s.Unlock() if s.state == Running { + if s.runCancel != nil { + s.runCancel() + } + if s.runDone != nil { + <-s.runDone + } s.server.Close() + s.runCancel = nil + s.runDone = nil } if err := os.RemoveAll(s.server.GetConfig().DataDir); err != nil { return err @@ -678,22 +745,117 @@ func restartTestCluster( // RunServer starts to run TestServer. func RunServer(server *TestServer) <-chan error { - resC := make(chan error) + resC := make(chan error, 1) go func() { resC <- server.Run() }() return resC } // RunServers starts to run multiple TestServer. func RunServers(servers []*TestServer) error { - res := make([]<-chan error, len(servers)) + runners := make([]testServerRunner, 0, len(servers)) + for _, server := range servers { + runners = append(runners, server) + } + return runTestServers(runners) +} + +type testServerRunner interface { + runWithStartSignal(chan<- struct{}) error + cancelRun() + Stop() error + State() int32 +} + +func runTestServers(servers []testServerRunner) error { + type runResult struct { + index int + err error + } + resC := make(chan runResult, len(servers)) for i, s := range servers { - res[i] = RunServer(s) + index := i + server := s + started := make(chan struct{}) + go func() { + resC <- runResult{index: index, err: server.runWithStartSignal(started)} + }() + // runWithStartSignal changes the state to Running before notifying started. + // Wait for that transition so an early failure cannot race with a server + // that has not entered run yet and therefore cannot be stopped. + <-started } - for _, c := range res { - if err := <-c; err != nil { - return errors.WithStack(err) + + var ( + primaryErr error + graceTimer *time.Timer + graceC <-chan time.Time + postCancelTimer *time.Timer + postCancelC <-chan time.Time + ) + remaining := len(servers) + completed := make([]bool, len(servers)) + handleResult := func(result runResult) { + remaining-- + completed[result.index] = true + if result.err != nil && primaryErr == nil { + primaryErr = result.err + graceTimer = time.NewTimer(runServersCleanupGracePeriod) + graceC = graceTimer.C } } + stopCompletedServers := func() { + for i, s := range servers { + if completed[i] && s.State() == Running { + _ = s.Stop() + } + } + } + for remaining > 0 { + select { + case result := <-resC: + handleResult(result) + case <-graceC: + graceC = nil + // Cancel without joining here: Stop waits for runDone and can block + // forever when startup does not observe its context. + for _, s := range servers { + if s.State() == Running { + s.cancelRun() + } + } + postCancelTimer = time.NewTimer(runServersCleanupGracePeriod) + postCancelC = postCancelTimer.C + case <-postCancelC: + postCancelC = nil + // Prefer results that became ready at the timeout boundary. + drainResults: + for remaining > 0 { + select { + case result := <-resC: + handleResult(result) + default: + break drainResults + } + } + if remaining > 0 { + stopCompletedServers() + return errors.Wrapf(errRunServersCleanupTimeout, + "server startup cleanup did not complete after cancellation; original error: %v", primaryErr) + } + } + } + if graceTimer != nil { + graceTimer.Stop() + } + if postCancelTimer != nil { + postCancelTimer.Stop() + } + if primaryErr != nil { + // Every Run call has returned, so closing these servers cannot race with + // partial startup. Do this before retry cleanup can remove their data dirs. + stopCompletedServers() + return errors.WithStack(primaryErr) + } return nil } diff --git a/tests/cluster_test.go b/tests/cluster_test.go index 5dcad634340..cf3346fcbad 100644 --- a/tests/cluster_test.go +++ b/tests/cluster_test.go @@ -16,7 +16,9 @@ package tests import ( "os" + "sync/atomic" "testing" + "time" "github.com/stretchr/testify/require" "go.uber.org/goleak" @@ -54,10 +56,166 @@ func TestClassifyInitialServersError(t *testing.T) { re.Equal(startServersRetryRecreate, classifyInitialServersError(errors.New("[PD:etcd:ErrStartEtcd]start etcd failed: listen tcp 127.0.0.1:2379: bind: address already in use"))) re.Equal(startServersRetryRecreate, classifyInitialServersError(errors.New("listen tcp 127.0.0.1:2379: bind: address already in use"))) re.Equal(startServersRetryRecreate, classifyInitialServersError(errors.New("Etcd cluster ID mismatch"))) + re.Equal(startServersNoRetry, classifyInitialServersError(errors.Wrap( + errRunServersCleanupTimeout, "listen tcp 127.0.0.1:2379: bind: address already in use"))) re.Equal(startServersNoRetry, classifyInitialServersError(errors.New("some other error"))) re.Equal(startServersNoRetry, classifyInitialServersError(nil)) } +type stubTestServer struct { + state atomic.Int32 + stopCh chan struct{} + runErr error + stopErr error + stopped atomic.Bool + canceled atomic.Bool + stopOnce atomic.Bool + finished atomic.Bool +} + +func newStubTestServer() *stubTestServer { + return &stubTestServer{ + stopCh: make(chan struct{}), + } +} + +func (s *stubTestServer) runWithStartSignal(started chan<- struct{}) error { + s.state.Store(Running) + if started != nil { + close(started) + } + if s.runErr != nil { + s.state.Store(Initial) + s.finished.Store(true) + return s.runErr + } + <-s.stopCh + s.finished.Store(true) + if s.stopErr != nil { + s.state.Store(Initial) + } + return s.stopErr +} + +func (s *stubTestServer) cancelRun() { + s.canceled.Store(true) + if s.stopOnce.CompareAndSwap(false, true) { + close(s.stopCh) + } +} + +func (s *stubTestServer) Stop() error { + if !s.state.CompareAndSwap(Running, Stop) { + return errors.New("server is not running") + } + s.stopped.Store(true) + if s.stopOnce.CompareAndSwap(false, true) { + close(s.stopCh) + } + return nil +} + +func (s *stubTestServer) State() int32 { + return s.state.Load() +} + +type uninterruptibleTestServer struct { + state atomic.Int32 + releaseCh chan struct{} + cancelCalled chan struct{} + cancelOnce atomic.Bool + releaseOnce atomic.Bool + finished atomic.Bool +} + +func newUninterruptibleTestServer() *uninterruptibleTestServer { + return &uninterruptibleTestServer{ + releaseCh: make(chan struct{}), + cancelCalled: make(chan struct{}), + } +} + +func (s *uninterruptibleTestServer) runWithStartSignal(started chan<- struct{}) error { + s.state.Store(Running) + close(started) + <-s.releaseCh + s.state.Store(Initial) + s.finished.Store(true) + return errors.New("start canceled") +} + +func (s *uninterruptibleTestServer) cancelRun() { + if s.cancelOnce.CompareAndSwap(false, true) { + close(s.cancelCalled) + } +} + +func (s *uninterruptibleTestServer) Stop() error { + // Model TestServer.Stop waiting for runDone while startup ignores cancellation. + <-s.releaseCh + s.state.Store(Stop) + return nil +} + +func (s *uninterruptibleTestServer) State() int32 { + return s.state.Load() +} + +func (s *uninterruptibleTestServer) release() { + if s.releaseOnce.CompareAndSwap(false, true) { + close(s.releaseCh) + } +} + +func TestRunServersCleanupIsBounded(t *testing.T) { + oldCleanupGracePeriod := runServersCleanupGracePeriod + runServersCleanupGracePeriod = time.Millisecond + t.Cleanup(func() { runServersCleanupGracePeriod = oldCleanupGracePeriod }) + + blocked := newUninterruptibleTestServer() + defer blocked.release() + failed := newStubTestServer() + failed.runErr = errors.New("start failed") + + done := make(chan error, 1) + go func() { + done <- runTestServers([]testServerRunner{blocked, failed}) + }() + + select { + case <-blocked.cancelCalled: + case <-time.After(time.Second): + t.Fatal("cleanup did not cancel the blocked server") + } + + select { + case err := <-done: + require.ErrorIs(t, err, errRunServersCleanupTimeout) + case <-time.After(50 * time.Millisecond): + t.Fatal("cleanup remained blocked after the grace period") + } + blocked.release() + require.Eventually(t, blocked.finished.Load, time.Second, time.Millisecond) +} + +func TestRunServersWaitsForInFlightRuns(t *testing.T) { + re := require.New(t) + oldCleanupGracePeriod := runServersCleanupGracePeriod + runServersCleanupGracePeriod = time.Millisecond + t.Cleanup(func() { runServersCleanupGracePeriod = oldCleanupGracePeriod }) + blockedServer := newStubTestServer() + blockedServer.stopErr = errors.New("start canceled") + failedServer := newStubTestServer() + failedServer.runErr = errors.New("start failed") + + // The cleanup error from the lower-index server must not replace the startup + // error that triggered cleanup. + err := runTestServers([]testServerRunner{blockedServer, failedServer}) + re.EqualError(err, "start failed") + re.True(blockedServer.canceled.Load()) + re.True(blockedServer.finished.Load()) +} + func TestRegenerateInitialServerURLsKeepsInitialClusterConsistent(t *testing.T) { t.Parallel() diff --git a/tests/integrations/mcs/keyspace/tso_keyspace_group_test.go b/tests/integrations/mcs/keyspace/tso_keyspace_group_test.go index 04cf95ad8f7..33acabcce89 100644 --- a/tests/integrations/mcs/keyspace/tso_keyspace_group_test.go +++ b/tests/integrations/mcs/keyspace/tso_keyspace_group_test.go @@ -39,8 +39,10 @@ import ( bs "github.com/tikv/pd/pkg/basicserver" "github.com/tikv/pd/pkg/keyspace" "github.com/tikv/pd/pkg/keyspace/constant" + tsoserver "github.com/tikv/pd/pkg/mcs/tso/server" mcs "github.com/tikv/pd/pkg/mcs/utils/constant" "github.com/tikv/pd/pkg/storage/endpoint" + "github.com/tikv/pd/pkg/utils/keypath" "github.com/tikv/pd/pkg/utils/tempurl" "github.com/tikv/pd/pkg/utils/testutil" "github.com/tikv/pd/pkg/utils/tsoutil" @@ -146,6 +148,27 @@ func (suite *keyspaceGroupTestSuite) closeAllTSONodesAndWait(re *require.Asserti }, testutil.WithWaitFor(10*time.Second), testutil.WithTickInterval(500*time.Millisecond)) } +func (suite *keyspaceGroupTestSuite) waitTSOKeyspaceGroupReady( + re *require.Assertions, + node *tsoserver.Server, + keyspaceID, keyspaceGroupID uint32, +) { + testutil.Eventually(re, func() bool { + resp, err := node.GetClient().Get(suite.ctx, keypath.KeyspaceGroupIDPath(keyspaceGroupID)) + if err != nil || len(resp.Kvs) == 0 { + return false + } + targetRevision := uint64(resp.Kvs[0].ModRevision) + _, kg, loadedGroupID, loadedRevision, err := + node.GetKeyspaceGroupManager().FindGroupByKeyspaceID(keyspaceID) + return err == nil && + kg != nil && + loadedGroupID == keyspaceGroupID && + loadedRevision >= targetRevision && + node.IsKeyspaceServingByGroup(keyspaceID, keyspaceGroupID) + }, testutil.WithWaitFor(30*time.Second), testutil.WithTickInterval(100*time.Millisecond)) +} + func (suite *keyspaceGroupTestSuite) TestAllocNodesUpdate() { re := suite.Require() // add three nodes. @@ -558,12 +581,13 @@ func (suite *keyspaceGroupTestSuite) trySetNodesForKeyspaceGroup(re *require.Ass // tsoTestSetup holds the setup information for TSO tests type tsoTestSetup struct { - nodes map[string]bs.Server - cleanups []func() - client pd.Client - initialTS uint64 - firstNodeAddr string - keyspaceID uint32 + nodes map[string]bs.Server + cleanups []func() + client pd.Client + initialTS uint64 + firstNodeAddr string + keyspaceID uint32 + keyspaceGroupID uint32 } // setupTSONodesAndClient creates TSO nodes, keyspace group, and returns initialized client @@ -649,12 +673,13 @@ func (suite *keyspaceGroupTestSuite) setupTSONodesAndClient(re *require.Assertio break } return &tsoTestSetup{ - nodes: nodes, - cleanups: cleanups, - client: client, - initialTS: initialTS, - firstNodeAddr: firstNodeAddr, - keyspaceID: keyspaceID, + nodes: nodes, + cleanups: cleanups, + client: client, + initialTS: initialTS, + firstNodeAddr: firstNodeAddr, + keyspaceID: keyspaceID, + keyspaceGroupID: keyspaceGroupID, } } @@ -736,6 +761,7 @@ func (suite *keyspaceGroupTestSuite) TestUpdateMemberWhenRecovery() { setup.cleanups = append(setup.cleanups, cleanup) nodes[newNode.GetAddr()] = newNode tests.WaitForPrimaryServing(re, map[string]bs.Server{newNode.GetAddr(): newNode}) + suite.waitTSOKeyspaceGroupReady(re, newNode, setup.keyspaceID, setup.keyspaceGroupID) // Step 7: Verify eventual recovery after node restart. // The in-flight GetTS may stay attached to stale discovery/metadata during diff --git a/tests/integrations/realcluster/etcd_key_test.go b/tests/integrations/realcluster/etcd_key_test.go index 77007b439bc..fb894fbd11e 100644 --- a/tests/integrations/realcluster/etcd_key_test.go +++ b/tests/integrations/realcluster/etcd_key_test.go @@ -78,6 +78,7 @@ var ( "/pd//scheduler_config/evict-stopping-store-scheduler", "/pd//timestamp", "/pd//tso/keyspace_groups/membership/", // ms + "/pd//tso/keyspace_groups/revision", // ms "/pd/cluster_id", } // The keys that prefix is `/ms`. @@ -95,6 +96,7 @@ var ( // These keys with `/pd` are only in `ms` mode. pdMSKeys = []string{ "/pd//tso/keyspace_groups/membership/", + "/pd//tso/keyspace_groups/revision", } ) diff --git a/tests/integrations/tso/consistency_test.go b/tests/integrations/tso/consistency_test.go index 766106993a1..b2a27ae5e0d 100644 --- a/tests/integrations/tso/consistency_test.go +++ b/tests/integrations/tso/consistency_test.go @@ -16,10 +16,13 @@ package tso import ( "context" + "net" + "net/url" "sync" "testing" "time" + "github.com/stretchr/testify/require" "github.com/stretchr/testify/suite" "google.golang.org/grpc" @@ -68,6 +71,68 @@ func TestMicroserviceTSOConsistencySuite(t *testing.T) { }) } +func TestRunInitialServersClosesStartingServersBeforeRetry(t *testing.T) { + re := require.New(t) + ctx, cancel := context.WithCancel(context.Background()) + cluster, err := tests.NewTestCluster(ctx, serverCount) + re.NoError(err) + t.Cleanup(func() { + cancel() + cluster.Destroy() + }) + + oldServer := cluster.GetServer("pd1") + oldClientURL, err := url.Parse(oldServer.GetConfig().ClientUrls) + re.NoError(err) + oldPeerURL, err := url.Parse(oldServer.GetConfig().PeerUrls) + re.NoError(err) + conflictURL, err := url.Parse(cluster.GetServer("pd3").GetConfig().ClientUrls) + re.NoError(err) + listener, err := net.Listen("tcp", conflictURL.Host) + re.NoError(err) + t.Cleanup(func() { re.NoError(listener.Close()) }) + + re.NoError(cluster.RunInitialServers()) + re.Equal(tests.Destroy, oldServer.State()) + for _, addr := range []string{oldClientURL.Host, oldPeerURL.Host} { + testutil.Eventually(re, func() bool { + conn, err := net.DialTimeout("tcp", addr, 100*time.Millisecond) + if conn != nil { + re.NoError(conn.Close()) + } + return err != nil + }, testutil.WithWaitFor(10*time.Second), testutil.WithTickInterval(100*time.Millisecond)) + } +} + +func TestRunFailureAfterEtcdStartClosesServer(t *testing.T) { + re := require.New(t) + ctx, cancel := context.WithCancel(context.Background()) + cluster, err := tests.NewTestCluster(ctx, 1) + re.NoError(err) + t.Cleanup(func() { + cancel() + cluster.Destroy() + }) + + const failpointName = "github.com/tikv/pd/server/failAfterStartEtcd" + re.NoError(failpoint.Enable(failpointName, "return(true)")) + t.Cleanup(func() { re.NoError(failpoint.Disable(failpointName)) }) + + testServer := cluster.GetServer("pd1") + clientURL, err := url.Parse(testServer.GetConfig().ClientUrls) + re.NoError(err) + err = cluster.RunInitialServers() + re.ErrorContains(err, "injected error after etcd startup") + re.NoError(testServer.Destroy()) + + conn, err := net.DialTimeout("tcp", clientURL.Host, time.Second) + re.Error(err) + if conn != nil { + re.NoError(conn.Close()) + } +} + func (suite *tsoConsistencyTestSuite) SetupSuite() { re := suite.Require()