diff --git a/client/go.mod b/client/go.mod index e99a4eca9af..0c4442cfdb7 100644 --- a/client/go.mod +++ b/client/go.mod @@ -10,7 +10,7 @@ require ( github.com/opentracing/opentracing-go v1.2.0 github.com/pingcap/errors v0.11.5-0.20211224045212-9687c2b0f87c github.com/pingcap/failpoint v0.0.0-20240528011301-b51a646c7c86 - github.com/pingcap/kvproto v0.0.0-20260903054228-107095f1d250 + github.com/pingcap/kvproto v0.0.0-20260903062353-65b4e27a438d github.com/pingcap/log v1.1.1-0.20221110025148-ca232912c9f3 github.com/prometheus/client_golang v1.20.5 github.com/prometheus/client_model v0.6.1 diff --git a/client/go.sum b/client/go.sum index b7e014b2e72..6a5b30aed15 100644 --- a/client/go.sum +++ b/client/go.sum @@ -53,8 +53,8 @@ github.com/pingcap/errors v0.11.5-0.20211224045212-9687c2b0f87c h1:xpW9bvK+HuuTm github.com/pingcap/errors v0.11.5-0.20211224045212-9687c2b0f87c/go.mod h1:X2r9ueLEUZgtx2cIogM0v4Zj5uvvzhuuiu7Pn8HzMPg= github.com/pingcap/failpoint v0.0.0-20240528011301-b51a646c7c86 h1:tdMsjOqUR7YXHoBitzdebTvOjs/swniBTOLy5XiMtuE= github.com/pingcap/failpoint v0.0.0-20240528011301-b51a646c7c86/go.mod h1:exzhVYca3WRtd6gclGNErRWb1qEgff3LYta0LvRmON4= -github.com/pingcap/kvproto v0.0.0-20260903054228-107095f1d250 h1:6yUryXKVbKpCNdZWL58/OcZj8NPLUA/xsJYXSbsD59w= -github.com/pingcap/kvproto v0.0.0-20260903054228-107095f1d250/go.mod h1:z6+aAHB7dBkA+LyinEX+48/ImRJ3jag0Hg0c7wkhEvE= +github.com/pingcap/kvproto v0.0.0-20260903062353-65b4e27a438d h1:KS3ekak/ljCj5xvkGqbwVLi2eL7B8GFSYzU9TOUUPPo= +github.com/pingcap/kvproto v0.0.0-20260903062353-65b4e27a438d/go.mod h1:z6+aAHB7dBkA+LyinEX+48/ImRJ3jag0Hg0c7wkhEvE= github.com/pingcap/log v1.1.1-0.20221110025148-ca232912c9f3 h1:HR/ylkkLmGdSSDaD8IDP+SZrdhV1Kibl9KrHxJ9eciw= github.com/pingcap/log v1.1.1-0.20221110025148-ca232912c9f3/go.mod h1:DWQW5jICDR7UJh4HtxXSM20Churx4CQL0fwL/SoOSA4= github.com/pkg/errors v0.8.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= diff --git a/errors.toml b/errors.toml index 33acad6fbfb..6c6fb3126e3 100644 --- a/errors.toml +++ b/errors.toml @@ -481,6 +481,11 @@ error = ''' trying to update GC safe point to a too large value that exceeds the txn safe point, current value: %v, given: %v, current txn safe point: %v ''' +["PD:gc:ErrGCStateWatcherSlowConsumer"] +error = ''' +gc state watcher is too slow +''' + ["PD:gc:ErrGlobalGCBarrierTSBehindTxnSafePoint"] error = ''' trying to set a global GC barrier on ts %d which is already behind the txn safe point %d of keyspace %s diff --git a/go.mod b/go.mod index 827bda4a3c0..6a011e018a0 100644 --- a/go.mod +++ b/go.mod @@ -35,7 +35,7 @@ require ( github.com/pingcap/errcode v0.3.0 github.com/pingcap/errors v0.11.5-0.20211224045212-9687c2b0f87c github.com/pingcap/failpoint v0.0.0-20240528011301-b51a646c7c86 - github.com/pingcap/kvproto v0.0.0-20260903054228-107095f1d250 + github.com/pingcap/kvproto v0.0.0-20260903062353-65b4e27a438d github.com/pingcap/log v1.1.1-0.20221110025148-ca232912c9f3 github.com/pingcap/metering_sdk v0.0.0-20260814062708-9e3b68cd9adf github.com/pingcap/sysutil v1.0.1-0.20230407040306-fb007c5aff21 diff --git a/go.sum b/go.sum index b74b2ade8cc..978e20fa993 100644 --- a/go.sum +++ b/go.sum @@ -490,8 +490,8 @@ github.com/pingcap/errors v0.11.5-0.20211224045212-9687c2b0f87c/go.mod h1:X2r9ue github.com/pingcap/failpoint v0.0.0-20240528011301-b51a646c7c86 h1:tdMsjOqUR7YXHoBitzdebTvOjs/swniBTOLy5XiMtuE= github.com/pingcap/failpoint v0.0.0-20240528011301-b51a646c7c86/go.mod h1:exzhVYca3WRtd6gclGNErRWb1qEgff3LYta0LvRmON4= github.com/pingcap/kvproto v0.0.0-20191211054548-3c6b38ea5107/go.mod h1:WWLmULLO7l8IOcQG+t+ItJ3fEcrL5FxF0Wu+HrMy26w= -github.com/pingcap/kvproto v0.0.0-20260903054228-107095f1d250 h1:6yUryXKVbKpCNdZWL58/OcZj8NPLUA/xsJYXSbsD59w= -github.com/pingcap/kvproto v0.0.0-20260903054228-107095f1d250/go.mod h1:z6+aAHB7dBkA+LyinEX+48/ImRJ3jag0Hg0c7wkhEvE= +github.com/pingcap/kvproto v0.0.0-20260903062353-65b4e27a438d h1:KS3ekak/ljCj5xvkGqbwVLi2eL7B8GFSYzU9TOUUPPo= +github.com/pingcap/kvproto v0.0.0-20260903062353-65b4e27a438d/go.mod h1:z6+aAHB7dBkA+LyinEX+48/ImRJ3jag0Hg0c7wkhEvE= github.com/pingcap/log v0.0.0-20210625125904-98ed8e2eb1c7/go.mod h1:8AanEdAHATuRurdGxZXBz0At+9avep+ub7U1AGYLIMM= github.com/pingcap/log v1.1.1-0.20221110025148-ca232912c9f3 h1:HR/ylkkLmGdSSDaD8IDP+SZrdhV1Kibl9KrHxJ9eciw= github.com/pingcap/log v1.1.1-0.20221110025148-ca232912c9f3/go.mod h1:DWQW5jICDR7UJh4HtxXSM20Churx4CQL0fwL/SoOSA4= diff --git a/pkg/errs/errno.go b/pkg/errs/errno.go index 5465e5dcf66..020f7dfc56b 100644 --- a/pkg/errs/errno.go +++ b/pkg/errs/errno.go @@ -548,10 +548,12 @@ var ( // GC errors var ( - ErrGCOnInvalidKeyspace = errors.Normalize("trying to manage GC in keyspace %v (id: %v) where keyspace level GC is not enabled", errors.RFCCodeText("PD:gc:ErrGCOnInvalidKeyspace")) - ErrDecreasingGCSafePoint = errors.Normalize("trying to update GC safe point to a smaller value, current value: %v, given: %v", errors.RFCCodeText("PD:gc:ErrDecreasingGCSafePoint")) - ErrGCSafePointExceedsTxnSafePoint = errors.Normalize("trying to update GC safe point to a too large value that exceeds the txn safe point, current value: %v, given: %v, current txn safe point: %v", errors.RFCCodeText("PD:gc:ErrGCSafePointExceedsTxnSafePoint")) - ErrDecreasingTxnSafePoint = errors.Normalize("trying to update txn safe point to a smaller value, current value: %v, given: %v", errors.RFCCodeText("PD:gc:ErrDecreasingTxnSafePoint")) + ErrGCOnInvalidKeyspace = errors.Normalize("trying to manage GC in keyspace %v (id: %v) where keyspace level GC is not enabled", errors.RFCCodeText("PD:gc:ErrGCOnInvalidKeyspace")) + ErrDecreasingGCSafePoint = errors.Normalize("trying to update GC safe point to a smaller value, current value: %v, given: %v", errors.RFCCodeText("PD:gc:ErrDecreasingGCSafePoint")) + ErrGCSafePointExceedsTxnSafePoint = errors.Normalize("trying to update GC safe point to a too large value that exceeds the txn safe point, current value: %v, given: %v, current txn safe point: %v", errors.RFCCodeText("PD:gc:ErrGCSafePointExceedsTxnSafePoint")) + ErrDecreasingTxnSafePoint = errors.Normalize("trying to update txn safe point to a smaller value, current value: %v, given: %v", errors.RFCCodeText("PD:gc:ErrDecreasingTxnSafePoint")) + // ErrGCStateWatcherSlowConsumer indicates that a watcher cannot keep up with live GC state changes. + ErrGCStateWatcherSlowConsumer = errors.Normalize("gc state watcher is too slow", errors.RFCCodeText("PD:gc:ErrGCStateWatcherSlowConsumer")) ErrGCBarrierTSBehindTxnSafePoint = errors.Normalize("trying to set a GC barrier on ts %d which is already behind the txn safe point %d", errors.RFCCodeText("PD:gc:ErrGCBarrierTSBehindTxnSafePoint")) ErrReservedGCBarrierID = errors.Normalize("trying to set a GC barrier with a barrier ID that is reserved: %v", errors.RFCCodeText("PD:gc:ErrReservedGCBarrierID")) ErrGlobalGCBarrierTSBehindTxnSafePoint = errors.Normalize("trying to set a global GC barrier on ts %d which is already behind the txn safe point %d of keyspace %s", errors.RFCCodeText("PD:gc:ErrGlobalGCBarrierTSBehindTxnSafePoint")) diff --git a/pkg/gc/enabled_keyspace_cache.go b/pkg/gc/enabled_keyspace_cache.go new file mode 100644 index 00000000000..04564a51d43 --- /dev/null +++ b/pkg/gc/enabled_keyspace_cache.go @@ -0,0 +1,391 @@ +// Copyright 2026 TiKV Project Authors. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package gc + +import ( + "context" + "fmt" + "slices" + "strconv" + "strings" + "sync" + "time" + + "github.com/gogo/protobuf/proto" + "go.etcd.io/etcd/api/v3/mvccpb" + clientv3 "go.etcd.io/etcd/client/v3" + "go.uber.org/zap" + + "github.com/pingcap/kvproto/pkg/keyspacepb" + "github.com/pingcap/log" + + "github.com/tikv/pd/pkg/keyspace" + "github.com/tikv/pd/pkg/utils/etcdutil" + "github.com/tikv/pd/pkg/utils/grpcutil" +) + +const ( + enabledKeyspacePageSize = 256 + enabledKeyspaceRequestTimeout = 5 * time.Second + enabledKeyspaceRetryDelay = time.Second + enabledKeyspaceMaxRetryDelay = 30 * time.Second + enabledKeyspaceLogInterval = 30 * time.Second + enabledKeyspaceWatchTimeout = 10 * time.Second +) + +// enabledKeyspace is a value copy of the metadata needed by GC initialization. +type enabledKeyspace struct { + id uint32 + gcManagementType string +} + +// enabledKeyspaceCache belongs to one leadership term. The published map and +// its revision always describe one complete, successfully applied snapshot. +type enabledKeyspaceCache struct { + termCtx context.Context + client *clientv3.Client + prefix string + watcherFactory func(*clientv3.Client) clientv3.Watcher + + mu sync.Mutex + entries map[uint32]enabledKeyspace + revision int64 + ready bool + changed chan struct{} +} + +func newEnabledKeyspaceCache(termCtx context.Context, client *clientv3.Client, prefix string) *enabledKeyspaceCache { + return &enabledKeyspaceCache{ + termCtx: termCtx, + client: client, + prefix: prefix, + watcherFactory: clientv3.NewWatcher, + changed: make(chan struct{}), + } +} + +// run blocks until the leadership term ends. A failed or compacted watch is +// followed by a complete reload, so no missing revision is silently skipped. +func (c *enabledKeyspaceCache) run() { + retryDelay := enabledKeyspaceRetryDelay + var lastLog time.Time + suppressedErrors := 0 + for c.termCtx.Err() == nil { + loadStartedAt := time.Now() + entries, revision, err := c.load() + phase := "load" + if err == nil { + if c.publish(entries, revision) { + log.Info("load enabled keyspace cache completed", + zap.String("prefix", c.prefix), + zap.Int64("revision", revision), + zap.Int("enabled-keyspace-count", len(entries)), + zap.Duration("cost", time.Since(loadStartedAt))) + } + watchStarted := time.Now() + err = c.watch(revision + 1) + phase = "watch" + if time.Since(watchStarted) >= enabledKeyspaceMaxRetryDelay { + retryDelay = enabledKeyspaceRetryDelay + } + } + if c.termCtx.Err() != nil { + return + } + if time.Since(lastLog) >= enabledKeyspaceLogInterval { + log.Warn("failed to synchronize enabled keyspace cache", + zap.String("prefix", c.prefix), + zap.String("phase", phase), + zap.Int64("revision", c.appliedRevision()), + zap.Duration("retry-delay", retryDelay), + zap.Int("suppressed-errors", suppressedErrors), + zap.Error(err)) + lastLog = time.Now() + suppressedErrors = 0 + } else { + suppressedErrors++ + } + timer := time.NewTimer(retryDelay) + select { + case <-c.termCtx.Done(): + timer.Stop() + return + case <-timer.C: + } + retryDelay = min(retryDelay*2, enabledKeyspaceMaxRetryDelay) + } +} + +func (c *enabledKeyspaceCache) appliedRevision() int64 { + c.mu.Lock() + defer c.mu.Unlock() + return c.revision +} + +func (c *enabledKeyspaceCache) waitReady(ctx context.Context) error { + for { + c.mu.Lock() + if err := ctx.Err(); err != nil { + c.mu.Unlock() + return err + } + if c.ready { + c.mu.Unlock() + return c.termCtx.Err() + } + changed := c.changed + c.mu.Unlock() + select { + case <-ctx.Done(): + return ctx.Err() + case <-c.termCtx.Done(): + return c.termCtx.Err() + case <-changed: + } + } +} + +func (c *enabledKeyspaceCache) snapshotAtLeast(ctx context.Context, revision int64) ([]enabledKeyspace, int64, error) { + for { + c.mu.Lock() + if err := c.termCtx.Err(); err != nil { + c.mu.Unlock() + return nil, 0, err + } + if err := ctx.Err(); err != nil { + c.mu.Unlock() + return nil, 0, err + } + if c.ready && c.revision >= revision { + result := make([]enabledKeyspace, 0, len(c.entries)) + for _, entry := range c.entries { + result = append(result, entry) + } + applied := c.revision + c.mu.Unlock() + slices.SortFunc(result, func(a, b enabledKeyspace) int { + switch { + case a.id < b.id: + return -1 + case a.id > b.id: + return 1 + default: + return 0 + } + }) + return result, applied, nil + } + changed := c.changed + c.mu.Unlock() + select { + case <-ctx.Done(): + return nil, 0, ctx.Err() + case <-c.termCtx.Done(): + return nil, 0, c.termCtx.Err() + case <-changed: + } + } +} + +func (c *enabledKeyspaceCache) publish(entries map[uint32]enabledKeyspace, revision int64) bool { + c.mu.Lock() + defer c.mu.Unlock() + if c.termCtx.Err() != nil { + return false + } + c.entries = entries + c.revision = revision + c.ready = true + close(c.changed) + c.changed = make(chan struct{}) + return true +} + +func (c *enabledKeyspaceCache) publishProgress(changes map[uint32]*enabledKeyspace, revision int64) { + c.mu.Lock() + defer c.mu.Unlock() + if c.termCtx.Err() != nil || revision <= c.revision { + return + } + for id, entry := range changes { + if entry == nil { + delete(c.entries, id) + } else { + c.entries[id] = *entry + } + } + c.revision = revision + close(c.changed) + c.changed = make(chan struct{}) +} + +// load reads each page at the first page's revision and returns only after the +// entire prefix has been decoded. No partial page is ever published. +func (c *enabledKeyspaceCache) load() (map[uint32]enabledKeyspace, int64, error) { + entries := make(map[uint32]enabledKeyspace) + start := c.prefix + end := clientv3.GetPrefixRangeEnd(c.prefix) + var revision int64 + for { + ctx, cancel := context.WithTimeout(c.termCtx, enabledKeyspaceRequestTimeout) + opts := []clientv3.OpOption{clientv3.WithRange(end), clientv3.WithLimit(enabledKeyspacePageSize)} + if revision != 0 { + opts = append(opts, clientv3.WithRev(revision)) + } + resp, err := c.client.Get(ctx, start, opts...) + cancel() + if err != nil { + return nil, 0, err + } + if revision == 0 { + revision = resp.Header.Revision + } + for _, kv := range resp.Kvs { + id, entry, enabled, err := c.decode(kv.Key, kv.Value) + if err != nil { + return nil, 0, err + } + if enabled { + entries[id] = entry + } + } + if !resp.More { + return entries, revision, nil + } + if len(resp.Kvs) == 0 { + return nil, 0, fmt.Errorf("empty keyspace metadata page at revision %d", revision) + } + start = string(resp.Kvs[len(resp.Kvs)-1].Key) + "\x00" + } +} + +func (c *enabledKeyspaceCache) decode(rawKey, rawValue []byte) (uint32, enabledKeyspace, bool, error) { + key := string(rawKey) + if !strings.HasPrefix(key, c.prefix) { + return 0, enabledKeyspace{}, false, fmt.Errorf("keyspace metadata key %q is outside prefix", key) + } + id64, err := strconv.ParseUint(strings.TrimPrefix(key, c.prefix), 10, 32) + if err != nil { + return 0, enabledKeyspace{}, false, fmt.Errorf("invalid keyspace metadata key %q: %w", key, err) + } + id := uint32(id64) + meta := &keyspacepb.KeyspaceMeta{} + if err := proto.Unmarshal(rawValue, meta); err != nil { + return 0, enabledKeyspace{}, false, fmt.Errorf("decode keyspace metadata %q: %w", key, err) + } + if meta.GetId() != id { + return 0, enabledKeyspace{}, false, fmt.Errorf("keyspace metadata %q contains ID %d", key, meta.GetId()) + } + return id, enabledKeyspace{id: id, gcManagementType: meta.Config[keyspace.GCManagementType]}, meta.State == keyspacepb.KeyspaceState_ENABLED, nil +} + +// watch updates a private working set. Progress notifications are ordered after +// prior events on the same stream, so they prove a complete applied revision. +// Event response headers alone may be ahead of the delivered prefix events. +func (c *enabledKeyspaceCache) watch(nextRevision int64) error { + watcher := c.watcherFactory(c.client) + defer watcher.Close() + watchCtx, cancel := context.WithCancel(clientv3.WithRequireLeader(c.termCtx)) + defer cancel() + done := make(chan struct{}) + go grpcutil.CheckStream(watchCtx, cancel, done) + watchCh := watcher.Watch(watchCtx, c.prefix, clientv3.WithPrefix(), clientv3.WithRev(nextRevision), clientv3.WithProgressNotify()) + done <- struct{}{} + if err := watchCtx.Err(); err != nil { + return fmt.Errorf("keyspace metadata watch creation failed: %w", err) + } + ticker := time.NewTicker(etcdutil.RequestProgressInterval) + defer ticker.Stop() + pending := make(map[uint32]*enabledKeyspace) + publishedRevision := nextRevision - 1 + pendingRevision := publishedRevision + lastProgress := time.Now() + for { + select { + case <-c.termCtx.Done(): + return c.termCtx.Err() + case <-ticker.C: + if time.Since(lastProgress) >= enabledKeyspaceWatchTimeout { + return fmt.Errorf("keyspace metadata watch made no progress for %s", enabledKeyspaceWatchTimeout) + } + ctx, cancel := context.WithTimeout(watchCtx, enabledKeyspaceRequestTimeout) + err := watcher.RequestProgress(ctx) + cancel() + if err != nil { + return err + } + case resp, ok := <-watchCh: + if !ok { + return fmt.Errorf("keyspace metadata watch closed") + } + if err := resp.Err(); err != nil { + return err + } + if resp.IsProgressNotify() { + if resp.Header.Revision < pendingRevision || resp.Header.Revision < publishedRevision || + (resp.Header.Revision == publishedRevision && len(pending) != 0) { + return fmt.Errorf("keyspace metadata watch progress %d precedes applied events at %d", resp.Header.Revision, pendingRevision) + } + c.publishProgress(pending, resp.Header.Revision) + publishedRevision = resp.Header.Revision + pendingRevision = publishedRevision + lastProgress = time.Now() + clear(pending) + continue + } + if len(resp.Events) == 0 { + continue + } + // Decode the complete clientv3 response before changing the working + // set. clientv3 merges fragmented watch responses before delivery. + type change struct { + id uint32 + entry enabledKeyspace + enabled bool + } + changes := make([]change, 0, len(resp.Events)) + for _, event := range resp.Events { + if event.Kv == nil || event.Kv.ModRevision <= publishedRevision { + return fmt.Errorf("keyspace metadata watch received stale or missing event at revision %d", publishedRevision) + } + pendingRevision = max(pendingRevision, event.Kv.ModRevision) + switch event.Type { + case mvccpb.PUT: + id, entry, enabled, err := c.decode(event.Kv.Key, event.Kv.Value) + if err != nil { + return err + } + changes = append(changes, change{id: id, entry: entry, enabled: enabled}) + case mvccpb.DELETE: + id64, err := strconv.ParseUint(strings.TrimPrefix(string(event.Kv.Key), c.prefix), 10, 32) + if err != nil { + return fmt.Errorf("invalid deleted keyspace metadata key %q: %w", event.Kv.Key, err) + } + changes = append(changes, change{id: uint32(id64)}) + default: + return fmt.Errorf("unexpected keyspace metadata event type %v", event.Type) + } + } + for _, change := range changes { + if change.enabled { + entry := change.entry + pending[change.id] = &entry + } else { + pending[change.id] = nil + } + } + } + } +} diff --git a/pkg/gc/enabled_keyspace_cache_test.go b/pkg/gc/enabled_keyspace_cache_test.go new file mode 100644 index 00000000000..a568021a985 --- /dev/null +++ b/pkg/gc/enabled_keyspace_cache_test.go @@ -0,0 +1,516 @@ +// Copyright 2026 TiKV Project Authors. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package gc + +import ( + "context" + "errors" + "fmt" + "net" + "sync" + "testing" + "time" + + "github.com/gogo/protobuf/proto" + "github.com/stretchr/testify/require" + pb "go.etcd.io/etcd/api/v3/etcdserverpb" + clientv3 "go.etcd.io/etcd/client/v3" + "google.golang.org/grpc" + + "github.com/pingcap/kvproto/pkg/keyspacepb" + + "github.com/tikv/pd/pkg/keyspace" + "github.com/tikv/pd/pkg/utils/etcdutil" +) + +const enabledKeyspaceTestPrefix = "/test/enabled-keyspaces/" + +type pauseAfterFirstPageKV struct { + clientv3.KV + firstPage chan<- int64 + release <-chan struct{} + once sync.Once +} + +func (kv *pauseAfterFirstPageKV) Get(ctx context.Context, key string, opts ...clientv3.OpOption) (*clientv3.GetResponse, error) { + resp, err := kv.KV.Get(ctx, key, opts...) + if err == nil { + kv.once.Do(func() { + kv.firstPage <- resp.Header.Revision + select { + case <-kv.release: + case <-ctx.Done(): + } + }) + } + return resp, err +} + +type pauseBeforeWatch struct { + clientv3.Watcher + started chan<- struct{} + release <-chan struct{} +} + +type signalWatchCreated struct { + clientv3.Watcher + created chan<- struct{} +} + +func (w *signalWatchCreated) Watch(ctx context.Context, key string, opts ...clientv3.OpOption) clientv3.WatchChan { + ch := w.Watcher.Watch(ctx, key, opts...) + w.created <- struct{}{} + return ch +} + +type neverCreateWatchServer struct { + pb.UnimplementedWatchServer + started chan struct{} +} + +func (s *neverCreateWatchServer) Watch(stream pb.Watch_WatchServer) error { + if _, err := stream.Recv(); err != nil { + return err + } + close(s.started) + <-stream.Context().Done() + return stream.Context().Err() +} + +func (w *pauseBeforeWatch) Watch(ctx context.Context, key string, opts ...clientv3.OpOption) clientv3.WatchChan { + w.started <- struct{}{} + select { + case <-w.release: + case <-ctx.Done(): + } + return w.Watcher.Watch(ctx, key, opts...) +} + +func putEnabledKeyspaceTestMeta(t *testing.T, client *clientv3.Client, id uint32, state keyspacepb.KeyspaceState, gcType string) int64 { + t.Helper() + value, err := proto.Marshal(&keyspacepb.KeyspaceMeta{ + Keyspace: &keyspacepb.KeyspaceMeta_Id{Id: id}, + State: state, + Config: map[string]string{keyspace.GCManagementType: gcType}, + }) + require.NoError(t, err) + resp, err := client.Put(context.Background(), fmt.Sprintf("%s%08d", enabledKeyspaceTestPrefix, id), string(value)) + require.NoError(t, err) + return resp.Header.Revision +} + +func startEnabledKeyspaceTestCache(t *testing.T, client *clientv3.Client) (*enabledKeyspaceCache, <-chan struct{}) { + t.Helper() + termCtx, cancel := context.WithCancel(context.Background()) + cache := newEnabledKeyspaceCache(termCtx, client, enabledKeyspaceTestPrefix) + done := make(chan struct{}) + go func() { + cache.run() + close(done) + }() + t.Cleanup(func() { + cancel() + select { + case <-done: + case <-time.After(5 * time.Second): + t.Error("cache did not stop after term cancellation") + } + }) + return cache, done +} + +func TestEnabledKeyspaceCacheEmptySnapshotAndProgress(t *testing.T) { + _, client, clean := etcdutil.NewTestEtcdCluster(t, 1, nil) + t.Cleanup(clean) + cache, _ := startEnabledKeyspaceTestCache(t, client) + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + require.NoError(t, cache.waitReady(ctx)) + initial, revision, err := cache.snapshotAtLeast(ctx, 0) + require.NoError(t, err) + require.Empty(t, initial) + require.Positive(t, revision) + + resp, err := client.Put(ctx, "/test/unrelated", "changed") + require.NoError(t, err) + list, applied, err := cache.snapshotAtLeast(ctx, resp.Header.Revision) + require.NoError(t, err) + require.Empty(t, list) + require.GreaterOrEqual(t, applied, resp.Header.Revision) +} + +func TestEnabledKeyspaceCacheAppliesMetadataChanges(t *testing.T) { + _, client, clean := etcdutil.NewTestEtcdCluster(t, 1, nil) + t.Cleanup(clean) + putEnabledKeyspaceTestMeta(t, client, 3, keyspacepb.KeyspaceState_ENABLED, keyspace.KeyspaceLevelGC) + putEnabledKeyspaceTestMeta(t, client, 2, keyspacepb.KeyspaceState_DISABLED, keyspace.UnifiedGC) + putEnabledKeyspaceTestMeta(t, client, 1, keyspacepb.KeyspaceState_ENABLED, keyspace.UnifiedGC) + cache, _ := startEnabledKeyspaceTestCache(t, client) + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + require.NoError(t, cache.waitReady(ctx)) + list, _, err := cache.snapshotAtLeast(ctx, 0) + require.NoError(t, err) + require.Equal(t, []enabledKeyspace{ + {id: 1, gcManagementType: keyspace.UnifiedGC}, + {id: 3, gcManagementType: keyspace.KeyspaceLevelGC}, + }, list) + + rev := putEnabledKeyspaceTestMeta(t, client, 2, keyspacepb.KeyspaceState_ENABLED, keyspace.KeyspaceLevelGC) + list, _, err = cache.snapshotAtLeast(ctx, rev) + require.NoError(t, err) + require.Equal(t, []enabledKeyspace{ + {id: 1, gcManagementType: keyspace.UnifiedGC}, + {id: 2, gcManagementType: keyspace.KeyspaceLevelGC}, + {id: 3, gcManagementType: keyspace.KeyspaceLevelGC}, + }, list) + + rev = putEnabledKeyspaceTestMeta(t, client, 1, keyspacepb.KeyspaceState_DISABLED, keyspace.UnifiedGC) + list, _, err = cache.snapshotAtLeast(ctx, rev) + require.NoError(t, err) + require.Equal(t, []enabledKeyspace{ + {id: 2, gcManagementType: keyspace.KeyspaceLevelGC}, + {id: 3, gcManagementType: keyspace.KeyspaceLevelGC}, + }, list) + + resp, err := client.Delete(ctx, fmt.Sprintf("%s%08d", enabledKeyspaceTestPrefix, 3)) + require.NoError(t, err) + list, _, err = cache.snapshotAtLeast(ctx, resp.Header.Revision) + require.NoError(t, err) + require.Equal(t, []enabledKeyspace{{id: 2, gcManagementType: keyspace.KeyspaceLevelGC}}, list) +} + +func TestEnabledKeyspaceCacheLoadsAllPagesAtOneRevision(t *testing.T) { + _, client, clean := etcdutil.NewTestEtcdCluster(t, 1, nil) + t.Cleanup(clean) + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + ops := make([]clientv3.Op, 0, enabledKeyspacePageSize+1) + var revision int64 + for id := uint32(1); id <= enabledKeyspacePageSize+1; id++ { + value, err := proto.Marshal(&keyspacepb.KeyspaceMeta{ + Keyspace: &keyspacepb.KeyspaceMeta_Id{Id: id}, + State: keyspacepb.KeyspaceState_ENABLED, + Config: map[string]string{keyspace.GCManagementType: keyspace.KeyspaceLevelGC}, + }) + require.NoError(t, err) + ops = append(ops, clientv3.OpPut(fmt.Sprintf("%s%08d", enabledKeyspaceTestPrefix, id), string(value))) + if len(ops) == 100 { + resp, err := client.Txn(ctx).Then(ops...).Commit() + require.NoError(t, err) + revision = resp.Header.Revision + ops = ops[:0] + } + } + resp, err := client.Txn(ctx).Then(ops...).Commit() + require.NoError(t, err) + revision = resp.Header.Revision + firstPage := make(chan int64, 1) + releasePage := make(chan struct{}) + clientWithPause := *client + clientWithPause.KV = &pauseAfterFirstPageKV{KV: client.KV, firstPage: firstPage, release: releasePage} + termCtx, stop := context.WithCancel(context.Background()) + defer stop() + cache := newEnabledKeyspaceCache(termCtx, &clientWithPause, enabledKeyspaceTestPrefix) + type loadedSnapshot struct { + entries map[uint32]enabledKeyspace + revision int64 + err error + } + loaded := make(chan loadedSnapshot, 1) + go func() { + entries, loadedRevision, err := cache.load() + loaded <- loadedSnapshot{entries: entries, revision: loadedRevision, err: err} + }() + select { + case firstRevision := <-firstPage: + require.Equal(t, revision, firstRevision) + case <-ctx.Done(): + t.Fatal("first metadata page was not read") + } + _, err = client.Delete(ctx, fmt.Sprintf("%s%08d", enabledKeyspaceTestPrefix, enabledKeyspacePageSize+1)) + require.NoError(t, err) + insertedRevision := putEnabledKeyspaceTestMeta(t, client, 300, keyspacepb.KeyspaceState_ENABLED, keyspace.UnifiedGC) + close(releasePage) + var initial loadedSnapshot + select { + case initial = <-loaded: + case <-ctx.Done(): + t.Fatal("fixed-revision metadata load did not finish") + } + require.NoError(t, initial.err) + require.Equal(t, revision, initial.revision) + require.Len(t, initial.entries, enabledKeyspacePageSize+1) + require.Contains(t, initial.entries, uint32(enabledKeyspacePageSize+1)) + require.NotContains(t, initial.entries, uint32(300)) + + cache.publish(initial.entries, initial.revision) + watchDone := make(chan error, 1) + go func() { watchDone <- cache.watch(initial.revision + 1) }() + list, applied, err := cache.snapshotAtLeast(ctx, insertedRevision) + require.NoError(t, err) + require.GreaterOrEqual(t, applied, insertedRevision) + require.Len(t, list, enabledKeyspacePageSize+1) + require.NotContains(t, list, enabledKeyspace{id: enabledKeyspacePageSize + 1, gcManagementType: keyspace.KeyspaceLevelGC}) + require.Equal(t, enabledKeyspace{id: 300, gcManagementType: keyspace.UnifiedGC}, list[len(list)-1]) + stop() + select { + case <-watchDone: + case <-ctx.Done(): + t.Fatal("metadata watch did not stop") + } +} + +func TestEnabledKeyspaceCacheRejectsMalformedMetadataUntilReload(t *testing.T) { + _, client, clean := etcdutil.NewTestEtcdCluster(t, 1, nil) + t.Cleanup(clean) + initial := putEnabledKeyspaceTestMeta(t, client, 1, keyspacepb.KeyspaceState_ENABLED, keyspace.KeyspaceLevelGC) + cache, _ := startEnabledKeyspaceTestCache(t, client) + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + require.NoError(t, cache.waitReady(ctx)) + + resp, err := client.Put(ctx, fmt.Sprintf("%s%08d", enabledKeyspaceTestPrefix, 2), "malformed protobuf") + require.NoError(t, err) + shortCtx, stop := context.WithTimeout(ctx, 300*time.Millisecond) + defer stop() + _, _, err = cache.snapshotAtLeast(shortCtx, resp.Header.Revision) + require.ErrorIs(t, err, context.DeadlineExceeded) + list, revision, err := cache.snapshotAtLeast(ctx, initial) + require.NoError(t, err) + require.Equal(t, initial, revision) + require.Equal(t, []enabledKeyspace{{id: 1, gcManagementType: keyspace.KeyspaceLevelGC}}, list) + + fixed := putEnabledKeyspaceTestMeta(t, client, 2, keyspacepb.KeyspaceState_ENABLED, keyspace.UnifiedGC) + list, revision, err = cache.snapshotAtLeast(ctx, fixed) + require.NoError(t, err) + require.GreaterOrEqual(t, revision, fixed) + require.Equal(t, []enabledKeyspace{ + {id: 1, gcManagementType: keyspace.KeyspaceLevelGC}, + {id: 2, gcManagementType: keyspace.UnifiedGC}, + }, list) +} + +func TestEnabledKeyspaceCacheReloadsAfterCompactedWatch(t *testing.T) { + _, client, clean := etcdutil.NewTestEtcdCluster(t, 1, nil) + t.Cleanup(clean) + first := putEnabledKeyspaceTestMeta(t, client, 1, keyspacepb.KeyspaceState_ENABLED, keyspace.KeyspaceLevelGC) + termCtx, cancel := context.WithCancel(context.Background()) + defer cancel() + cache := newEnabledKeyspaceCache(termCtx, client, enabledKeyspaceTestPrefix) + watchStarting := make(chan struct{}, 1) + releaseWatch := make(chan struct{}) + firstWatch := true + cache.watcherFactory = func(client *clientv3.Client) clientv3.Watcher { + watcher := clientv3.NewWatcher(client) + if !firstWatch { + return watcher + } + firstWatch = false + return &pauseBeforeWatch{Watcher: watcher, started: watchStarting, release: releaseWatch} + } + done := make(chan struct{}) + go func() { + cache.run() + close(done) + }() + ctx, stop := context.WithTimeout(context.Background(), 10*time.Second) + defer stop() + select { + case <-watchStarting: + case <-ctx.Done(): + t.Fatal("initial cache load did not reach watch startup") + } + list, revision, err := cache.snapshotAtLeast(ctx, first) + require.NoError(t, err) + require.Equal(t, first, revision) + require.Equal(t, []enabledKeyspace{{id: 1, gcManagementType: keyspace.KeyspaceLevelGC}}, list) + _, err = client.Delete(ctx, fmt.Sprintf("%s%08d", enabledKeyspaceTestPrefix, 1)) + require.NoError(t, err) + latest := putEnabledKeyspaceTestMeta(t, client, 2, keyspacepb.KeyspaceState_ENABLED, keyspace.UnifiedGC) + _, err = client.Compact(ctx, latest, clientv3.WithCompactPhysical()) + require.NoError(t, err) + close(releaseWatch) + list, applied, err := cache.snapshotAtLeast(ctx, latest) + require.NoError(t, err) + require.GreaterOrEqual(t, applied, latest) + require.Equal(t, []enabledKeyspace{{id: 2, gcManagementType: keyspace.UnifiedGC}}, list) + cancel() + select { + case <-done: + case <-ctx.Done(): + t.Fatal("cache run did not stop") + } +} + +func TestEnabledKeyspaceCacheTermCancellationUnblocksWaiters(t *testing.T) { + _, client, clean := etcdutil.NewTestEtcdCluster(t, 1, nil) + t.Cleanup(clean) + termCtx, cancel := context.WithCancel(context.Background()) + cache := newEnabledKeyspaceCache(termCtx, client, enabledKeyspaceTestPrefix) + done := make(chan struct{}) + go func() { + cache.run() + close(done) + }() + ctx, stop := context.WithTimeout(context.Background(), 10*time.Second) + defer stop() + require.NoError(t, cache.waitReady(ctx)) + _, revision, err := cache.snapshotAtLeast(ctx, 0) + require.NoError(t, err) + waitResult := make(chan error, 1) + go func() { + _, _, err := cache.snapshotAtLeast(ctx, revision+100) + waitResult <- err + }() + cancel() + select { + case err := <-waitResult: + require.True(t, errors.Is(err, context.Canceled), "waiting snapshot error: %v", err) + case <-ctx.Done(): + t.Fatal("snapshot did not stop after term cancellation") + } + select { + case <-done: + case <-ctx.Done(): + t.Fatal("cache run did not stop after term cancellation") + } +} + +func TestEnabledKeyspaceCacheWatchCreationTimesOut(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + server := grpc.NewServer() + backend := &neverCreateWatchServer{started: make(chan struct{})} + pb.RegisterWatchServer(server, backend) + go func() { _ = server.Serve(listener) }() + defer server.Stop() + client, err := clientv3.New(clientv3.Config{Endpoints: []string{listener.Addr().String()}, DialTimeout: time.Second}) + require.NoError(t, err) + defer client.Close() + termCtx, cancel := context.WithCancel(context.Background()) + defer cancel() + cache := newEnabledKeyspaceCache(termCtx, client, enabledKeyspaceTestPrefix) + done := make(chan error, 1) + go func() { done <- cache.watch(1) }() + select { + case <-backend.started: + case <-time.After(3 * time.Second): + t.Fatal("watch create request did not reach server") + } + select { + case err := <-done: + require.Error(t, err) + require.NoError(t, termCtx.Err()) + case <-time.After(6 * time.Second): + cancel() + <-done + t.Fatal("watch creation did not time out") + } +} + +func TestEnabledKeyspaceCacheReloadsAfterWatchCreationTimeout(t *testing.T) { + _, client, clean := etcdutil.NewTestEtcdCluster(t, 1, nil) + t.Cleanup(clean) + first := putEnabledKeyspaceTestMeta(t, client, 1, keyspacepb.KeyspaceState_ENABLED, keyspace.KeyspaceLevelGC) + termCtx, cancel := context.WithCancel(context.Background()) + cache := newEnabledKeyspaceCache(termCtx, client, enabledKeyspaceTestPrefix) + watchStarted := make(chan struct{}, 1) + releaseWatch := make(chan struct{}) + firstWatch := true + cache.watcherFactory = func(client *clientv3.Client) clientv3.Watcher { + watcher := clientv3.NewWatcher(client) + if !firstWatch { + return watcher + } + firstWatch = false + return &pauseBeforeWatch{Watcher: watcher, started: watchStarted, release: releaseWatch} + } + done := make(chan struct{}) + go func() { + cache.run() + close(done) + }() + defer func() { + cancel() + select { + case <-done: + case <-time.After(5 * time.Second): + t.Error("cache did not stop") + } + }() + ctx, stop := context.WithTimeout(context.Background(), 8*time.Second) + defer stop() + select { + case <-watchStarted: + case <-ctx.Done(): + t.Fatal("initial cache load did not reach watch creation") + } + initial, revision, err := cache.snapshotAtLeast(ctx, first) + require.NoError(t, err) + require.Equal(t, first, revision) + require.Equal(t, []enabledKeyspace{{id: 1, gcManagementType: keyspace.KeyspaceLevelGC}}, initial) + latest := putEnabledKeyspaceTestMeta(t, client, 2, keyspacepb.KeyspaceState_ENABLED, keyspace.UnifiedGC) + list, applied, err := cache.snapshotAtLeast(ctx, latest) + require.NoError(t, err) + require.GreaterOrEqual(t, applied, latest) + require.Equal(t, []enabledKeyspace{ + {id: 1, gcManagementType: keyspace.KeyspaceLevelGC}, + {id: 2, gcManagementType: keyspace.UnifiedGC}, + }, list) +} + +func TestEnabledKeyspaceCacheWatchSurvivesCreationTimeout(t *testing.T) { + _, client, clean := etcdutil.NewTestEtcdCluster(t, 1, nil) + t.Cleanup(clean) + termCtx, stop := context.WithCancel(context.Background()) + cache := newEnabledKeyspaceCache(termCtx, client, enabledKeyspaceTestPrefix) + created := make(chan struct{}, 1) + cache.watcherFactory = func(client *clientv3.Client) clientv3.Watcher { + return &signalWatchCreated{Watcher: clientv3.NewWatcher(client), created: created} + } + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + entries, revision, err := cache.load() + require.NoError(t, err) + require.True(t, cache.publish(entries, revision)) + done := make(chan error, 1) + go func() { done <- cache.watch(revision + 1) }() + defer func() { + stop() + select { + case <-done: + case <-time.After(5 * time.Second): + t.Error("watch did not stop") + } + }() + select { + case <-created: + case <-ctx.Done(): + t.Fatal("watch was not created") + } + select { + case err := <-done: + t.Fatalf("watch ended after creation: %v", err) + case <-time.After(4 * time.Second): + } + latest := putEnabledKeyspaceTestMeta(t, client, 3, keyspacepb.KeyspaceState_ENABLED, keyspace.UnifiedGC) + list, applied, err := cache.snapshotAtLeast(ctx, latest) + require.NoError(t, err) + require.GreaterOrEqual(t, applied, latest) + require.Equal(t, []enabledKeyspace{{id: 3, gcManagementType: keyspace.UnifiedGC}}, list) +} diff --git a/pkg/gc/gc_state_manager.go b/pkg/gc/gc_state_manager.go index aa25bbd9fda..40897157ca4 100644 --- a/pkg/gc/gc_state_manager.go +++ b/pkg/gc/gc_state_manager.go @@ -23,6 +23,7 @@ import ( "sync/atomic" "time" + clientv3 "go.etcd.io/etcd/client/v3" "go.uber.org/zap" "github.com/pingcap/failpoint" @@ -198,15 +199,18 @@ type GCStateManager struct { // A read/write - update cache procedure must be done while holding the outer mutex `GCStateManager.mu`. // A read-only operation can be done on gcStateCache directly without locking `GCStateManager.mu`. - gcStateCache *gcStateCache + gcStateCache *gcStateCache + etcdClient *clientv3.Client + enabledKeyspaces *enabledKeyspaceCache + cancelEnabledKeyspaces context.CancelFunc allKeyspacesGCStatesSingleFlight *syncutil.OrderedSingleFlight[map[uint32]GCState] allKeyspacesGCStatesExcludeGCBarriersSingleFlight *syncutil.OrderedSingleFlight[map[uint32]GCState] - // Note that nodeLeadership is a counter instead of a bool. Theoretically, it's possible that an - // OnNodeBecomesFollower invocation of the previous lease is later than the OnNodeBecomesLeader call of the new - // lease during PD leader changes. Making this a counter helps in guaranteeing the eventual consistency. - nodeLeadership atomic.Int32 + watchers map[uint64]*GCStateWatcher + nextWatcherID uint64 + nextLeadershipGeneration uint64 + activeLeadershipGeneration atomic.Uint64 } // NewGCStateManager creates a GCStateManager of GC and services. @@ -216,6 +220,7 @@ func NewGCStateManager(store endpoint.GCStateProvider, cfg config.PDServerConfig cfg: cfg, keyspaceManager: keyspaceManager, gcStateCache: newGCStateCache(), + watchers: make(map[uint64]*GCStateWatcher), allKeyspacesGCStatesSingleFlight: syncutil.NewOrderedSingleFlight[map[uint32]GCState](), allKeyspacesGCStatesExcludeGCBarriersSingleFlight: syncutil.NewOrderedSingleFlight[map[uint32]GCState](), } @@ -226,6 +231,14 @@ func NewGCStateManager(store endpoint.GCStateProvider, cfg config.PDServerConfig return m } +// SetEtcdClient supplies the client used by the leader-local keyspace index. +// It must be called before the manager's first leadership generation starts. +func (m *GCStateManager) SetEtcdClient(client *clientv3.Client) { + m.mu.Lock() + defer m.mu.Unlock() + m.etcdClient = client +} + type keyspaceNameKeyType struct{} var keyspaceNameKey = keyspaceNameKeyType{} @@ -244,37 +257,57 @@ func getKeyspaceNameFromCtx(ctx context.Context) string { return "" } -// OnNodeBecomesLeader marks the current PD node as leader for GC state watches. -func (m *GCStateManager) OnNodeBecomesLeader() { +// OnNodeBecomesLeader starts a local leadership generation and returns its teardown function. +func (m *GCStateManager) OnNodeBecomesLeader() func() { m.mu.Lock() - defer m.mu.Unlock() - - m.nodeLeadership.Add(1) - - // Also trigger cache invalidation even when transitioning from follower to leader, as a protection against - // potential inconsistent cache state left from the last leadership. + // Disable lock-free cache reads throughout the reset, including when + // replacing an active leadership generation. + m.activeLeadershipGeneration.Store(0) + if m.cancelEnabledKeyspaces != nil { + m.cancelEnabledKeyspaces() + } + m.nextLeadershipGeneration++ + generation := m.nextLeadershipGeneration + m.terminateAllGCStateWatchersLocked(errs.ErrNotLeader, watcherTerminationLeaderLost) + failpoint.InjectCall("beforeLeaderGCStateCacheReset") m.gcStateCache.clearAll() + m.enabledKeyspaces = nil + m.cancelEnabledKeyspaces = nil + if m.etcdClient != nil { + termCtx, cancel := context.WithCancel(context.Background()) + m.cancelEnabledKeyspaces = cancel + m.enabledKeyspaces = newEnabledKeyspaceCache(termCtx, m.etcdClient, keypath.KeyspaceMetaPrefix()) + } + enabledKeyspaces := m.enabledKeyspaces m.barrierMetrics.clearMetrics() productionBarrierMetrics.current.Store(m.barrierMetrics) -} - -// OnNodeBecomesFollower marks the current PD node as follower and closes all existing GC state watches. -func (m *GCStateManager) OnNodeBecomesFollower() { - m.mu.Lock() - defer m.mu.Unlock() - - m.nodeLeadership.Add(-1) + m.activeLeadershipGeneration.Store(generation) + m.mu.Unlock() + if enabledKeyspaces != nil { + go enabledKeyspaces.run() + } - // Invalidate the cache. - m.gcStateCache.clearAll() - m.barrierMetrics.clearMetrics() - if !m.nodeIsLeader() { + return func() { + m.mu.Lock() + defer m.mu.Unlock() + if m.activeLeadershipGeneration.Load() != generation { + return + } + m.activeLeadershipGeneration.Store(0) + if m.cancelEnabledKeyspaces != nil { + m.cancelEnabledKeyspaces() + m.cancelEnabledKeyspaces = nil + m.enabledKeyspaces = nil + } + m.terminateAllGCStateWatchersLocked(errs.ErrNotLeader, watcherTerminationLeaderLost) + m.gcStateCache.clearAll() + m.barrierMetrics.clearMetrics() productionBarrierMetrics.current.CompareAndSwap(m.barrierMetrics, nil) } } func (m *GCStateManager) nodeIsLeader() bool { - return m.nodeLeadership.Load() > 0 + return m.activeLeadershipGeneration.Load() != 0 } // redirectKeyspace checks the given keyspaceID, and returns the actual keyspaceID to operate on. @@ -406,6 +439,14 @@ func (m *GCStateManager) advanceGCSafePointImpl(ctx context.Context, keyspaceID TxnSafePoint: txnSafePoint, GCSafePoint: newGCSafePoint, }) + if newGCSafePoint != oldGCSafePoint { + m.publishGCStateChangeLocked(NewGCStateUpsert(GCState{ + KeyspaceID: keyspaceID, + IsKeyspaceLevel: keyspaceID != constant.NullKeyspaceID, + TxnSafePoint: txnSafePoint, + GCSafePoint: newGCSafePoint, + })) + } if newGCSafePoint != oldGCSafePoint { log.Info("advanced GC safe point", @@ -593,6 +634,14 @@ func (m *GCStateManager) advanceTxnSafePointImpl(ctx context.Context, keyspaceID TxnSafePoint: newTxnSafePoint, GCSafePoint: gcSafePoint, }) + if newTxnSafePoint != oldTxnSafePoint { + m.publishGCStateChangeLocked(NewGCStateUpsert(GCState{ + KeyspaceID: keyspaceID, + IsKeyspaceLevel: keyspaceID != constant.NullKeyspaceID, + TxnSafePoint: newTxnSafePoint, + GCSafePoint: gcSafePoint, + })) + } blockerDesc := "" simulatedServiceID := "" diff --git a/pkg/gc/gc_state_manager_test.go b/pkg/gc/gc_state_manager_test.go index ddcdd6af714..ab54d4dac44 100644 --- a/pkg/gc/gc_state_manager_test.go +++ b/pkg/gc/gc_state_manager_test.go @@ -98,6 +98,8 @@ type newGCStateManagerForTestOptions struct { serverNodes int etcdServerCfgModifier func(cfg *embed.Config) etcdClientCfgModifier etcdutil.CreateEtcdClientOpt + useEnabledKeyspaceCache bool + beforeLeader func(*clientv3.Client) } func (opt *newGCStateManagerForTestOptions) generateKeyspacesByCount(count int) { @@ -144,6 +146,9 @@ func newGCStateManagerForTest(t testing.TB, opt newGCStateManagerForTestOptions) kgm := keyspace.NewKeyspaceGroupManager(ctx, s, client) keyspaceManager := keyspace.NewKeyspaceManager(ctx, s, mockcluster.NewCluster(ctx, config.NewPersistOptions(cfg)), allocator, &config.KeyspaceConfig{}, kgm, nil) gcStateManager = NewGCStateManager(s.GetGCStateProvider(), cfg.PDServerCfg, keyspaceManager) + if opt.useEnabledKeyspaceCache { + gcStateManager.SetEtcdClient(client) + } t.Cleanup(gcStateManager.CloseBarrierMetrics) err = kgm.Bootstrap(ctx) @@ -210,7 +215,15 @@ func newGCStateManagerForTest(t testing.TB, opt newGCStateManagerForTestOptions) } } - gcStateManager.OnNodeBecomesLeader() + if opt.beforeLeader != nil { + opt.beforeLeader(client) + } + stopGCStateManager := gcStateManager.OnNodeBecomesLeader() + originalClean := clean + clean = func() { + stopGCStateManager() + originalClean() + } return s, s.GetGCStateProvider(), gcStateManager, clean, cancel } @@ -259,10 +272,321 @@ type gcStateCacheAccessCounterSnapshot struct { func (s *gcStateManagerTestSuite) ensureMarkedLeader() { if !s.manager.nodeIsLeader() { - s.manager.OnNodeBecomesLeader() + stopGCStateManager := s.manager.OnNodeBecomesLeader() + s.T().Cleanup(stopGCStateManager) + } +} + +func (s *gcStateManagerTestSuite) TestGCStateWatchRequiresActiveLeadership() { + follower := NewGCStateManager(s.provider, s.manager.cfg, s.manager.keyspaceManager) + _, err := follower.WatchGCStates(context.Background(), true) + s.Require().ErrorIs(err, errs.ErrNotLeader) +} + +func (s *gcStateManagerTestSuite) TestGCStateWatchLeadershipGeneration() { + re := s.Require() + stopFirst := s.manager.OnNodeBecomesLeader() + first, err := s.manager.WatchGCStates(context.Background(), true) + re.NoError(err) + + stopSecond := s.manager.OnNodeBecomesLeader() + re.ErrorIs(first.Err(), errs.ErrNotLeader) + second, err := s.manager.WatchGCStates(context.Background(), true) + re.NoError(err) + + stopFirst() + re.NoError(second.Err()) + stopSecond() + re.ErrorIs(second.Err(), errs.ErrNotLeader) +} + +func TestGCStateLeadershipClearsCacheBeforeEnablingReads(t *testing.T) { + for _, replacingLeader := range []bool{false, true} { + name := "follower-promotion" + if replacingLeader { + name = "leader-replacement" + } + t.Run(name, func(t *testing.T) { + re := require.New(t) + _, client, clean := etcdutil.NewTestEtcdCluster(t, 1, nil) + defer clean() + provider := endpoint.NewStorageEndpoint(kv.NewEtcdKVBase(client), nil).GetGCStateProvider() + manager := NewGCStateManager(provider, config.PDServerConfig{}, nil) + defer manager.CloseBarrierMetrics() + if replacingLeader { + stop := manager.OnNodeBecomesLeader() + defer stop() + } + + writeSafePoint := func(safePoint uint64) { + re.NoError(provider.RunInGCStateTransaction(func(wb *endpoint.GCStateWriteBatch) error { + return wb.SetGCSafePoint(constant.NullKeyspaceID, safePoint) + })) + } + writeSafePoint(100) + // Reads can populate the cache even while this manager is a follower. + state, err := manager.GetGCState(constant.NullKeyspaceID, true) + re.NoError(err) + re.Equal(uint64(100), state.GCSafePoint) + // Another leader advances storage after the cached read. + writeSafePoint(200) + + resetReached := make(chan struct{}) + releaseReset := make(chan struct{}) + release := sync.OnceFunc(func() { close(releaseReset) }) + const hook = "github.com/tikv/pd/pkg/gc/beforeLeaderGCStateCacheReset" + re.NoError(failpoint.EnableCall(hook, func() { + close(resetReached) + <-releaseReset + })) + defer func() { re.NoError(failpoint.Disable(hook)) }() + leaderReady := make(chan struct{}) + var stopLeader func() + go func() { + stopLeader = manager.OnNodeBecomesLeader() + close(leaderReady) + }() + defer func() { + release() + select { + case <-leaderReady: + stopLeader() + case <-time.After(5 * time.Second): + t.Error("leadership initialization did not stop") + } + }() + select { + case <-resetReached: + case <-time.After(5 * time.Second): + t.Fatal("leadership initialization did not reach the cache reset") + } + + // Cache reads must remain disabled while the stale snapshot is present. + safePoint, err := manager.CompatibleLoadGCSafePoint(constant.NullKeyspaceID) + re.NoError(err) + re.Equal(uint64(200), safePoint) + release() + select { + case <-leaderReady: + case <-time.After(5 * time.Second): + t.Fatal("leadership initialization did not finish") + } + re.True(manager.nodeIsLeader()) + safePoint, err = manager.CompatibleLoadGCSafePoint(constant.NullKeyspaceID) + re.NoError(err) + re.Equal(uint64(200), safePoint) + }) } } +func (s *gcStateManagerTestSuite) TestGCStateWatchLoadsInitialStatesIncrementally() { + re := s.Require() + w, err := s.manager.registerGCStateWatcher(context.Background(), false, gcStateWatchConfig{ + initialBatchSize: 2, + initChannelCapacity: 1, + liveChannelCapacity: 1, + }) + re.NoError(err) + defer w.Close() + + want := make(map[uint32]struct{}, len(s.keyspacePresets.all)) + for _, keyspaceID := range s.keyspacePresets.all { + want[keyspaceID] = struct{}{} + } + got := make(map[uint32]struct{}, len(want)) + var batchSizes []int + for len(got) < len(want) { + changes, err := w.RecvBatch(2) + re.NoError(err) + batchSizes = append(batchSizes, len(changes)) + for _, change := range changes { + state := mustUpsert(s.T(), change) + re.Nil(state.GCBarriers) + got[state.KeyspaceID] = struct{}{} + } + } + re.Equal([]int{2, 2, 1}, batchSizes) + re.Equal(want, got) +} + +func (s *gcStateManagerTestSuite) TestGCStateWatchSkipsInitialLoading() { + re := s.Require() + w, err := s.manager.WatchGCStates(context.Background(), true) + re.NoError(err) + defer w.Close() + re.True(w.initDone) + select { + case batch := <-w.initCh: + re.Fail("initial channel was used", "batch: %+v", batch) + default: + } +} + +func (s *gcStateManagerTestSuite) TestGCStateWatchCanceledBeforeInitialLoaderDoesNotIterate() { + re := s.Require() + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + var iterationStarted atomic.Bool + re.NoError(failpoint.EnableCall("github.com/tikv/pd/pkg/gc/onGetAllKeyspacesGCStatesStart", func() { + iterationStarted.Store(true) + })) + defer func() { re.NoError(failpoint.Disable("github.com/tikv/pd/pkg/gc/onGetAllKeyspacesGCStatesStart")) }() + re.NoError(failpoint.EnableCall("github.com/tikv/pd/pkg/gc/watchGCStatesRegistered", cancel)) + defer func() { re.NoError(failpoint.Disable("github.com/tikv/pd/pkg/gc/watchGCStatesRegistered")) }() + + // Suppress the automatic loader so this test can model a loader goroutine that + // starts only after registration has synchronously canceled its watcher. + w, err := s.manager.registerGCStateWatcher(ctx, true, gcStateWatchConfig{initialBatchSize: 1, liveChannelCapacity: 1}) + re.NoError(err) + defer w.Close() + re.ErrorIs(w.Err(), context.Canceled) + + s.manager.loadInitialGCStates(w, 1) + re.False(iterationStarted.Load()) +} + +func (s *gcStateManagerTestSuite) TestGCStateWatchLiveSuppressesPausedInitial() { + re := s.Require() + const keyspaceID = uint32(2) + _, err := s.manager.AdvanceTxnSafePoint(keyspaceID, 10, time.Now()) + re.NoError(err) + + reached := make(chan struct{}) + release := make(chan struct{}) + var reachedOnce, releaseOnce sync.Once + releaseLoader := func() { releaseOnce.Do(func() { close(release) }) } + re.NoError(failpoint.EnableCall("github.com/tikv/pd/pkg/gc/watchGCStatesInitialStateLoaded", func(id uint32) { + if id == keyspaceID { + reachedOnce.Do(func() { close(reached) }) + <-release + } + })) + defer func() { re.NoError(failpoint.Disable("github.com/tikv/pd/pkg/gc/watchGCStatesInitialStateLoaded")) }() + defer releaseLoader() + + w, err := s.manager.registerGCStateWatcher(context.Background(), false, gcStateWatchConfig{initialBatchSize: 1, initChannelCapacity: 16, liveChannelCapacity: 4}) + re.NoError(err) + defer w.Close() + select { + case <-reached: + case <-time.After(5 * time.Second): + re.FailNow("initial loader did not reach keyspace 2") + } + + _, err = s.manager.AdvanceTxnSafePoint(keyspaceID, 20, time.Now()) + re.NoError(err) + for { + changes, err := w.RecvBatch(1) + re.NoError(err) + state, ok := changes[0].Upsert() + if ok && state.KeyspaceID == keyspaceID && state.TxnSafePoint == 20 { + break + } + } + releaseLoader() + + re.Eventually(func() bool { + for { + change, ok, err := w.receiveOne(false) + re.NoError(err) + if !ok { + return w.initDone + } + state, upsert := change.Upsert() + re.False(upsert && state.KeyspaceID == keyspaceID && state.TxnSafePoint == 10) + } + }, 5*time.Second, 10*time.Millisecond) +} + +func (s *gcStateManagerTestSuite) TestGCStateWatchInitialFailureTerminatesWatcher() { + re := s.Require() + const errorMessage = "injected initial watch failure" + re.NoError(failpoint.Enable("github.com/tikv/pd/pkg/gc/iterateAllKeyspacesGCStatesError", fmt.Sprintf(`return(%q)`, errorMessage))) + defer func() { re.NoError(failpoint.Disable("github.com/tikv/pd/pkg/gc/iterateAllKeyspacesGCStatesError")) }() + + w, err := s.manager.WatchGCStates(context.Background(), false) + re.NoError(err) + _, err = w.RecvBatch(1) + re.ErrorContains(err, errorMessage) + re.Eventually(func() bool { + s.manager.mu.RLock() + defer s.manager.mu.RUnlock() + _, ok := s.manager.watchers[w.id] + return !ok + }, 5*time.Second, 10*time.Millisecond) +} + +func (s *gcStateManagerTestSuite) TestGCStateWatchFullInitChannelDoesNotHoldManagerMutex() { + re := s.Require() + stop := s.manager.OnNodeBecomesLeader() + w, err := s.manager.registerGCStateWatcher(context.Background(), false, gcStateWatchConfig{ + initialBatchSize: 1, + initChannelCapacity: 1, + liveChannelCapacity: 1, + }) + re.NoError(err) + re.Eventually(func() bool { return len(w.initCh) == cap(w.initCh) }, 5*time.Second, 10*time.Millisecond) + + mutationDone := make(chan error, 1) + go func() { + _, err := s.manager.AdvanceTxnSafePoint(2, 1, time.Now()) + mutationDone <- err + }() + select { + case err := <-mutationDone: + re.NoError(err) + case <-time.After(5 * time.Second): + re.FailNow("manager mutation blocked behind the initial state loader") + } + + teardownDone := make(chan struct{}) + go func() { + stop() + close(teardownDone) + }() + select { + case <-teardownDone: + case <-time.After(5 * time.Second): + re.FailNow("leadership teardown blocked behind the initial state loader") + } + w.Close() + re.Eventually(func() bool { + s.manager.mu.RLock() + defer s.manager.mu.RUnlock() + _, ok := s.manager.watchers[w.id] + return !ok + }, 5*time.Second, 10*time.Millisecond) +} + +func (s *gcStateManagerTestSuite) TestGCStateWatchConcurrentCloseIsIdempotent() { + re := s.Require() + stop := s.manager.OnNodeBecomesLeader() + w, err := s.manager.WatchGCStates(context.Background(), false) + re.NoError(err) + + var wg sync.WaitGroup + wg.Add(3) + go func() { + defer wg.Done() + w.Close() + }() + go func() { + defer wg.Done() + s.manager.terminateGCStateWatcher(w, errors.New("initial load failed"), watcherTerminationInitError) + }() + go func() { + defer wg.Done() + stop() + }() + wg.Wait() + re.Error(w.Err()) + s.manager.mu.RLock() + _, ok := s.manager.watchers[w.id] + s.manager.mu.RUnlock() + re.False(ok) +} + func (s *gcStateManagerTestSuite) trackGCStateCacheAccessCounters() *gcStateCacheAccessCounters { tracker := &gcStateCacheAccessCounters{} failpointName := "github.com/tikv/pd/pkg/gc/getGCStateCacheAccess" @@ -571,9 +895,9 @@ func (s *gcStateManagerTestSuite) TestCompatibleUpdateGCSafePointSequentiallyWit return wb.SetGCSafePoint(keyspaceID, 101) }) re.NoError(err) - oldLeadership := s.manager.nodeLeadership.Load() - s.manager.nodeLeadership.Store(0) - defer s.manager.nodeLeadership.Store(oldLeadership) + oldLeadership := s.manager.activeLeadershipGeneration.Load() + s.manager.activeLeadershipGeneration.Store(0) + defer s.manager.activeLeadershipGeneration.Store(oldLeadership) gcSafePoint, err = s.manager.CompatibleLoadGCSafePoint(keyspaceID) re.NoError(err) @@ -2078,7 +2402,8 @@ func (s *gcStateManagerTestSuite) TestGetGCStateWithGlobalGCBarriersRejectsRevis s.manager.keyspaceManager, ) s.T().Cleanup(otherManager.CloseBarrierMetrics) - otherManager.OnNodeBecomesLeader() + stopOtherManager := otherManager.OnNodeBecomesLeader() + defer stopOtherManager() _, err = otherManager.SetGlobalGCBarrier( ctx, "snapshot", diff --git a/pkg/gc/gc_state_watcher.go b/pkg/gc/gc_state_watcher.go new file mode 100644 index 00000000000..793db7f9759 --- /dev/null +++ b/pkg/gc/gc_state_watcher.go @@ -0,0 +1,512 @@ +// Copyright 2026 TiKV Project Authors. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package gc + +import ( + "context" + "fmt" + "time" + + "go.uber.org/zap" + + "github.com/pingcap/errors" + "github.com/pingcap/failpoint" + "github.com/pingcap/log" + + "github.com/tikv/pd/pkg/errs" + "github.com/tikv/pd/pkg/keyspace" + "github.com/tikv/pd/pkg/keyspace/constant" + "github.com/tikv/pd/pkg/utils/keypath" +) + +type gcStateChangeKind uint8 + +const ( + gcStateChangeUnknown gcStateChangeKind = iota + gcStateChangeUpsert + gcStateChangeRemoved +) + +// GCStateChange describes one effective GC state change for a keyspace scope. +// nolint:revive // Keep GC in the name to match the established GCState domain API. +type GCStateChange struct { + kind gcStateChangeKind + upsert GCState + removedKeyspaceID uint32 +} + +// NewGCStateUpsert creates a change containing the complete effective GC state. +func NewGCStateUpsert(state GCState) GCStateChange { + state.GCBarriers = nil + return GCStateChange{kind: gcStateChangeUpsert, upsert: state} +} + +// NewGCStateRemoved creates a change that removes a keyspace scope. +func NewGCStateRemoved(keyspaceID uint32) GCStateChange { + return GCStateChange{kind: gcStateChangeRemoved, removedKeyspaceID: keyspaceID} +} + +// Upsert returns the effective GC state when the change is an upsert. +func (c GCStateChange) Upsert() (GCState, bool) { + return c.upsert, c.kind == gcStateChangeUpsert +} + +// RemovedKeyspaceID returns the removed keyspace ID when the change is a removal. +func (c GCStateChange) RemovedKeyspaceID() (uint32, bool) { + return c.removedKeyspaceID, c.kind == gcStateChangeRemoved +} + +// KeyspaceID returns the keyspace scope changed by this value. +func (c GCStateChange) KeyspaceID() (uint32, bool) { + if state, ok := c.Upsert(); ok { + return state.KeyspaceID, true + } + return c.RemovedKeyspaceID() +} + +const ( + defaultGCStateWatchInitialBatchSize = 1024 + defaultGCStateWatchInitChannelCapacity = 1 + defaultGCStateWatchLiveChannelCapacity = 1024 + gcStateWatchMetadataWaitTimeout = 5 * time.Minute +) + +type gcStateWatchConfig struct { + initialBatchSize int + initChannelCapacity int + liveChannelCapacity int +} + +type gcStateWatcherTerminationReason string + +const ( + watcherTerminationClientCancel gcStateWatcherTerminationReason = "client_cancel" + watcherTerminationLeaderLost gcStateWatcherTerminationReason = "leader_lost" + watcherTerminationSlowConsumer gcStateWatcherTerminationReason = "slow_consumer" + watcherTerminationInitError gcStateWatcherTerminationReason = "init_error" +) + +// GCStateWatcher merges ordered initial batches and live GC state changes for one stream. +// The initial scan is not globally atomic and may interleave with live delivery. For each +// keyspace, the merge prevents initial and live delivery from regressing to an older state. +// +// A watcher supports one receiving goroutine. Close may be called concurrently with +// receiving and with manager-owned lifecycle operations. +// nolint:revive // Keep GC in the name to match the established GCState domain API. +type GCStateWatcher struct { + ctx context.Context + cancel context.CancelCauseFunc + manager *GCStateManager + id uint64 + initCh chan []GCStateChange + liveCh chan GCStateChange + initDone bool + pendingInit []GCStateChange + pendingLiveCount int + dirtyDuringInit map[uint32]struct{} + enabledKeyspaces *enabledKeyspaceCache +} + +func newGCStateWatcher(parent context.Context, cfg gcStateWatchConfig, skipLoadingInitial bool) *GCStateWatcher { + ctx, cancel := context.WithCancelCause(parent) + watcher := &GCStateWatcher{ + ctx: ctx, + cancel: cancel, + initCh: make(chan []GCStateChange, cfg.initChannelCapacity), + liveCh: make(chan GCStateChange, cfg.liveChannelCapacity), + initDone: skipLoadingInitial, + } + if !skipLoadingInitial { + watcher.dirtyDuringInit = make(map[uint32]struct{}) + } + return watcher +} + +func (w *GCStateWatcher) receiveOne(block bool) (GCStateChange, bool, error) { + for { + if err := w.Err(); err != nil { + return GCStateChange{}, false, err + } + + if w.pendingLiveCount > 0 { + // This watcher has one receiver, so every change counted when the + // initial batch was acquired is still queued until we consume it. + change := <-w.liveCh + w.pendingLiveCount-- + if keyspaceID, valid := change.KeyspaceID(); valid { + w.dirtyDuringInit[keyspaceID] = struct{}{} + } + return change, true, nil + } + + for len(w.pendingInit) > 0 { + change := w.pendingInit[0] + w.pendingInit = w.pendingInit[1:] + keyspaceID, ok := change.KeyspaceID() + if ok { + if _, dirty := w.dirtyDuringInit[keyspaceID]; dirty { + // Registration precedes the initial scan, so a post-registration live v2 may + // race with initial v1. If v1 is consumed first, delivery is v1 then v2; if + // v2 is consumed first, this later v1 is suppressed and delivery is v2 only. + continue + } + } + return change, true, nil + } + w.pendingInit = nil + + if w.initDone { + if block { + select { + case <-w.ctx.Done(): + return GCStateChange{}, false, w.Err() + case change := <-w.liveCh: + return change, true, nil + } + } + select { + case <-w.ctx.Done(): + return GCStateChange{}, false, w.Err() + case change := <-w.liveCh: + return change, true, nil + default: + return GCStateChange{}, false, nil + } + } + + var ( + change GCStateChange + batch []GCStateChange + ok bool + ) + if block { + select { + case <-w.ctx.Done(): + return GCStateChange{}, false, w.Err() + case change = <-w.liveCh: + if keyspaceID, valid := change.KeyspaceID(); valid { + // Marking live scopes dirty preserves the alternate v2-only order when + // initial v1 has not yet been delivered. + w.dirtyDuringInit[keyspaceID] = struct{}{} + } + return change, true, nil + case batch, ok = <-w.initCh: + } + } else { + select { + case <-w.ctx.Done(): + return GCStateChange{}, false, w.Err() + case change = <-w.liveCh: + if keyspaceID, valid := change.KeyspaceID(); valid { + w.dirtyDuringInit[keyspaceID] = struct{}{} + } + return change, true, nil + case batch, ok = <-w.initCh: + default: + return GCStateChange{}, false, nil + } + } + + if !ok { + w.initCh = nil + w.initDone = true + w.dirtyDuringInit = nil + continue + } + failpoint.InjectCall("watchGCStatesInitialBatchReceived") + w.pendingInit = batch + // Snapshot only after acquiring the initial batch. Registration precedes + // its scan, and mutations publish under the manager mutex before the next + // mutation can update the cache. Thus any live state older than this batch's + // initial state is already queued or consumed. The initial state's own live + // publication may still follow its cache store, but cannot cause regression. + // Drain this FIFO prefix first and suppress initial scopes it makes dirty. + // Keep the remaining count across RecvBatch calls; later arrivals must not + // extend the prefix and indefinitely postpone unrelated initial states. + w.pendingLiveCount = len(w.liveCh) + } +} + +// Done returns a channel that is closed when the watcher terminates. +func (w *GCStateWatcher) Done() <-chan struct{} { + return w.ctx.Done() +} + +// Err returns the first cause that terminated the watcher. +func (w *GCStateWatcher) Err() error { + return context.Cause(w.ctx) +} + +// RecvBatch waits for one visible change and opportunistically collects up to maxChanges. +func (w *GCStateWatcher) RecvBatch(maxChanges int) ([]GCStateChange, error) { + if maxChanges <= 0 { + panic("GCStateWatcher.RecvBatch requires a positive maximum") + } + first, ok, err := w.receiveOne(true) + if err != nil { + return nil, err + } + if !ok { + panic("blocking watcher receive returned no result") + } + result := []GCStateChange{first} + for len(result) < maxChanges { + change, ok, err := w.receiveOne(false) + if err != nil { + return nil, err + } + if !ok { + break + } + result = append(result, change) + } + if err := w.Err(); err != nil { + return nil, err + } + return result, nil +} + +// Close stops the watcher and removes it from its manager. +func (w *GCStateWatcher) Close() { + if w.manager == nil { + w.cancel(context.Canceled) + return + } + w.manager.terminateGCStateWatcher(w, context.Canceled, watcherTerminationClientCancel) +} + +// WatchGCStates registers a watcher in the current local leadership generation. +func (m *GCStateManager) WatchGCStates(ctx context.Context, skipLoadingInitial bool) (*GCStateWatcher, error) { + return m.registerGCStateWatcher(ctx, skipLoadingInitial, gcStateWatchConfig{ + initialBatchSize: defaultGCStateWatchInitialBatchSize, + initChannelCapacity: defaultGCStateWatchInitChannelCapacity, + liveChannelCapacity: defaultGCStateWatchLiveChannelCapacity, + }) +} + +func (m *GCStateManager) registerGCStateWatcher( + ctx context.Context, + skipLoadingInitial bool, + cfg gcStateWatchConfig, +) (*GCStateWatcher, error) { + watcher := newGCStateWatcher(ctx, cfg, skipLoadingInitial) + var cache *enabledKeyspaceCache + var generation uint64 + if !skipLoadingInitial { + m.mu.RLock() + generation = m.activeLeadershipGeneration.Load() + cache = m.enabledKeyspaces + m.mu.RUnlock() + if generation == 0 { + watcher.cancel(errs.ErrNotLeader) + return nil, errs.ErrNotLeader + } + if cache != nil { + failpoint.InjectCall("watchGCStatesBeforeCacheReady") + waitCtx, cancel := context.WithTimeout(ctx, gcStateWatchMetadataWaitTimeout) + err := cache.waitReady(waitCtx) + cancel() + if err != nil { + if ctx.Err() == nil && cache.termCtx.Err() != nil { + err = errs.ErrNotLeader + } + watcher.cancel(err) + return nil, err + } + } + } + + m.mu.Lock() + if m.activeLeadershipGeneration.Load() == 0 || (generation != 0 && m.activeLeadershipGeneration.Load() != generation) { + m.mu.Unlock() + watcher.cancel(errs.ErrNotLeader) + return nil, errs.ErrNotLeader + } + m.nextWatcherID++ + watcher.manager = m + watcher.id = m.nextWatcherID + watcher.enabledKeyspaces = cache + m.watchers[watcher.id] = watcher + gcStateWatcherGauge.Inc() + m.mu.Unlock() + + failpoint.InjectCall("watchGCStatesRegistered") + if !skipLoadingInitial { + go m.loadInitialGCStates(watcher, cfg.initialBatchSize) + } + return watcher, nil +} + +func (m *GCStateManager) loadInitialGCStates(watcher *GCStateWatcher, batchSize int) { + if watcher.Err() != nil { + return + } + + batch := make([]GCStateChange, 0, batchSize) + stopped := false + flush := func() bool { + if len(batch) == 0 { + return true + } + ready := batch + batch = make([]GCStateChange, 0, batchSize) + select { + case watcher.initCh <- ready: + return true + case <-watcher.ctx.Done(): + return false + } + } + + addState := func(state GCState) { + if stopped { + return + } + failpoint.InjectCall("watchGCStatesInitialStateLoaded", state.KeyspaceID) + if watcher.Err() != nil { + stopped = true + return + } + batch = append(batch, NewGCStateUpsert(state)) + if len(batch) == batchSize { + stopped = !flush() + } + } + var err error + if watcher.enabledKeyspaces != nil { + err = m.iterateEnabledKeyspacesGCStates(watcher.ctx, watcher.enabledKeyspaces, addState) + } else { + err = m.iterateAllKeyspacesGCStates(watcher.ctx, true, func(uint32) bool { return true }, addState, nil) + } + + if stopped || watcher.Err() != nil { + return + } + if err != nil { + m.terminateGCStateWatcher(watcher, errors.Annotate(err, "load initial GC states"), watcherTerminationInitError) + return + } + if !flush() { + return + } + close(watcher.initCh) +} + +func (m *GCStateManager) iterateEnabledKeyspacesGCStates( + ctx context.Context, + cache *enabledKeyspaceCache, + cb func(GCState), +) error { + // The default Get is linearizable. This exact key supplies only the global + // revision; metadata membership comes from the shared index. + probeCtx, cancel := context.WithTimeout(ctx, enabledKeyspaceRequestTimeout) + resp, err := cache.client.Get(probeCtx, keypath.KeyspaceMetaPrefix()) + cancel() + if err != nil { + return fmt.Errorf("probe keyspace metadata revision: %w", err) + } + failpoint.InjectCall("watchGCStatesTargetRevisionProbed", resp.Header.Revision) + waitCtx, cancel := context.WithTimeout(ctx, gcStateWatchMetadataWaitTimeout) + entries, _, err := cache.snapshotAtLeast(waitCtx, resp.Header.Revision) + cancel() + if err != nil { + return fmt.Errorf("wait for keyspace metadata revision %d: %w", resp.Header.Revision, err) + } + + nullState, err := m.getGCStateImpl(constant.NullKeyspaceID, true) + if err != nil { + return err + } + cb(nullState) + for _, entry := range entries { + if err := ctx.Err(); err != nil { + return err + } + if entry.gcManagementType != keyspace.KeyspaceLevelGC { + cb(GCState{KeyspaceID: entry.id, IsKeyspaceLevel: false}) + continue + } + state, err := m.getGCStateImpl(entry.id, true) + if err != nil { + return err + } + cb(state) + } + return nil +} + +func (m *GCStateManager) terminateGCStateWatcher( + watcher *GCStateWatcher, + cause error, + reason gcStateWatcherTerminationReason, +) { + m.mu.Lock() + defer m.mu.Unlock() + m.terminateGCStateWatcherLocked(watcher, cause, reason) +} + +func (m *GCStateManager) publishGCStateChangeLocked(change GCStateChange) { + for _, watcher := range m.watchers { + select { + case watcher.liveCh <- change: + default: + log.Warn("GC state watcher is too slow", + zap.Uint64("watcher-id", watcher.id), + zap.Int("capacity", cap(watcher.liveCh)), + zap.Int("queue-length", len(watcher.liveCh))) + m.terminateGCStateWatcherLocked(watcher, errs.ErrGCStateWatcherSlowConsumer, watcherTerminationSlowConsumer) + } + } + // TODO: Publish keyspace metadata upserts and removals through this same serialized path when an authoritative GC-leader-owned lifecycle hook exists. +} + +func (m *GCStateManager) terminateGCStateWatcherLocked( + watcher *GCStateWatcher, + cause error, + reason gcStateWatcherTerminationReason, +) { + registered, ok := m.watchers[watcher.id] + if !ok || registered != watcher { + return + } + delete(m.watchers, watcher.id) + gcStateWatcherGauge.Dec() + recordGCStateWatcherTerminationMetrics(reason) + watcher.cancel(cause) +} + +func recordGCStateWatcherTerminationMetrics(reason gcStateWatcherTerminationReason) { + switch reason { + case watcherTerminationClientCancel: + gcStateWatcherTerminationClientCancelCounter.Inc() + case watcherTerminationLeaderLost: + gcStateWatcherTerminationLeaderLostCounter.Inc() + case watcherTerminationSlowConsumer: + gcStateWatcherTerminationSlowConsumerCounter.Inc() + case watcherTerminationInitError: + gcStateWatcherTerminationInitErrorCounter.Inc() + default: + panic("unknown GC state watcher termination reason") + } +} + +func (m *GCStateManager) terminateAllGCStateWatchersLocked( + cause error, + reason gcStateWatcherTerminationReason, +) { + for _, watcher := range m.watchers { + m.terminateGCStateWatcherLocked(watcher, cause, reason) + } +} diff --git a/pkg/gc/gc_state_watcher_test.go b/pkg/gc/gc_state_watcher_test.go new file mode 100644 index 00000000000..861aa8b9b7a --- /dev/null +++ b/pkg/gc/gc_state_watcher_test.go @@ -0,0 +1,761 @@ +// Copyright 2026 TiKV Project Authors. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package gc + +import ( + "context" + "errors" + "fmt" + "math" + "sync" + "testing" + "time" + + promtestutil "github.com/prometheus/client_golang/prometheus/testutil" + "github.com/stretchr/testify/require" + clientv3 "go.etcd.io/etcd/client/v3" + + "github.com/pingcap/failpoint" + + "github.com/tikv/pd/pkg/errs" + "github.com/tikv/pd/pkg/keyspace" + "github.com/tikv/pd/pkg/utils/keypath" +) + +func TestGCStateWatcherInitialThenLive(t *testing.T) { + w := newGCStateWatcher(context.Background(), gcStateWatchConfig{initChannelCapacity: 1, liveChannelCapacity: 2}, false) + w.initCh <- []GCStateChange{NewGCStateUpsert(GCState{KeyspaceID: 7, TxnSafePoint: 1})} + + got, err := w.RecvBatch(1) + require.NoError(t, err) + require.Equal(t, uint64(1), mustUpsert(t, got[0]).TxnSafePoint) + + w.liveCh <- NewGCStateUpsert(GCState{KeyspaceID: 7, TxnSafePoint: 2}) + got, err = w.RecvBatch(1) + require.NoError(t, err) + require.Equal(t, uint64(2), mustUpsert(t, got[0]).TxnSafePoint) +} + +func TestGCStateWatcherLiveSuppressesOlderInitial(t *testing.T) { + w := newGCStateWatcher(context.Background(), gcStateWatchConfig{initChannelCapacity: 1, liveChannelCapacity: 2}, false) + w.liveCh <- NewGCStateUpsert(GCState{KeyspaceID: 7, TxnSafePoint: 2}) + + got, err := w.RecvBatch(1) + require.NoError(t, err) + require.Equal(t, uint64(2), mustUpsert(t, got[0]).TxnSafePoint) + + w.initCh <- []GCStateChange{NewGCStateUpsert(GCState{KeyspaceID: 7, TxnSafePoint: 1})} + close(w.initCh) + _, ok, err := w.receiveOne(false) + require.NoError(t, err) + require.False(t, ok) + require.True(t, w.initDone) +} + +func TestGCStateWatcherQueuedLivePrecedesNewerInitial(t *testing.T) { + for _, maxChanges := range []int{1, 2, 4} { + t.Run(fmt.Sprintf("batch-size-%d", maxChanges), func(t *testing.T) { + w := newGCStateWatcher(context.Background(), gcStateWatchConfig{initChannelCapacity: 2, liveChannelCapacity: 2}, false) + t.Cleanup(w.Close) + w.initCh <- []GCStateChange{ + NewGCStateUpsert(GCState{KeyspaceID: 7, TxnSafePoint: 20}), + NewGCStateUpsert(GCState{KeyspaceID: 8, TxnSafePoint: 1}), + } + w.initCh <- []GCStateChange{ + NewGCStateUpsert(GCState{KeyspaceID: 7, TxnSafePoint: 20}), + NewGCStateUpsert(GCState{KeyspaceID: 9, TxnSafePoint: 1}), + } + close(w.initCh) + queueLiveWhenInitialReceived(t, w, + NewGCStateUpsert(GCState{KeyspaceID: 7, TxnSafePoint: 10}), + NewGCStateUpsert(GCState{KeyspaceID: 7, TxnSafePoint: 20}), + ) + + want := []GCState{ + {KeyspaceID: 7, TxnSafePoint: 10}, + {KeyspaceID: 7, TxnSafePoint: 20}, + {KeyspaceID: 8, TxnSafePoint: 1}, + {KeyspaceID: 9, TxnSafePoint: 1}, + } + for offset := 0; offset < len(want); { + got, err := w.RecvBatch(maxChanges) + require.NoError(t, err) + require.Len(t, got, min(maxChanges, len(want)-offset)) + for _, change := range got { + require.Equal(t, want[offset], mustUpsert(t, change)) + offset++ + } + } + _, ok, err := w.receiveOne(false) + require.NoError(t, err) + require.False(t, ok, "initial duplicates must remain suppressed through the closed channel's buffered batches") + }) + } +} + +func TestGCStateWatcherLaterLiveDoesNotPostponePendingInitial(t *testing.T) { + w := newGCStateWatcher(context.Background(), gcStateWatchConfig{initChannelCapacity: 1, liveChannelCapacity: 4}, false) + t.Cleanup(w.Close) + w.initCh <- []GCStateChange{ + NewGCStateUpsert(GCState{KeyspaceID: 7, TxnSafePoint: 20}), + NewGCStateUpsert(GCState{KeyspaceID: 8, TxnSafePoint: 1}), + NewGCStateUpsert(GCState{KeyspaceID: 9, TxnSafePoint: 1}), + } + close(w.initCh) + queueLiveWhenInitialReceived(t, w, + NewGCStateUpsert(GCState{KeyspaceID: 7, TxnSafePoint: 10}), + NewGCStateUpsert(GCState{KeyspaceID: 7, TxnSafePoint: 20}), + ) + + for i, want := range []GCState{ + {KeyspaceID: 7, TxnSafePoint: 10}, + {KeyspaceID: 7, TxnSafePoint: 20}, + {KeyspaceID: 8, TxnSafePoint: 1}, + {KeyspaceID: 9, TxnSafePoint: 1}, + } { + got, err := w.RecvBatch(1) + require.NoError(t, err) + require.Len(t, got, 1) + require.Equal(t, want, mustUpsert(t, got[0])) + // Keep the live queue nonempty after the first receive. These arrivals + // must not postpone the unrelated states in the acquired initial batch. + w.liveCh <- NewGCStateUpsert(GCState{KeyspaceID: 7, TxnSafePoint: uint64(30 + i*10)}) + } + got, err := w.RecvBatch(4) + require.NoError(t, err) + require.Len(t, got, 4) + for i, want := range []uint64{30, 40, 50, 60} { + require.Equal(t, GCState{KeyspaceID: 7, TxnSafePoint: want}, mustUpsert(t, got[i])) + } +} + +func TestGCStateWatcherQueuedRemovalSuppressesInitial(t *testing.T) { + w := newGCStateWatcher(context.Background(), gcStateWatchConfig{initChannelCapacity: 1, liveChannelCapacity: 2}, false) + t.Cleanup(w.Close) + w.initCh <- []GCStateChange{ + NewGCStateUpsert(GCState{KeyspaceID: 7, TxnSafePoint: 10}), + NewGCStateUpsert(GCState{KeyspaceID: 8, TxnSafePoint: 1}), + } + close(w.initCh) + queueLiveWhenInitialReceived(t, w, NewGCStateRemoved(7)) + + got, err := w.RecvBatch(3) + require.NoError(t, err) + require.Len(t, got, 2) + removed, ok := got[0].RemovedKeyspaceID() + require.True(t, ok) + require.Equal(t, uint32(7), removed) + require.Equal(t, GCState{KeyspaceID: 8, TxnSafePoint: 1}, mustUpsert(t, got[1])) +} + +func TestGCStateWatcherCancellationDiscardsPendingLivePrefix(t *testing.T) { + w := newGCStateWatcher(context.Background(), gcStateWatchConfig{initChannelCapacity: 1, liveChannelCapacity: 2}, false) + t.Cleanup(w.Close) + w.initCh <- []GCStateChange{ + NewGCStateUpsert(GCState{KeyspaceID: 7, TxnSafePoint: 20}), + NewGCStateUpsert(GCState{KeyspaceID: 8, TxnSafePoint: 1}), + } + close(w.initCh) + queueLiveWhenInitialReceived(t, w, + NewGCStateUpsert(GCState{KeyspaceID: 7, TxnSafePoint: 10}), + NewGCStateUpsert(GCState{KeyspaceID: 7, TxnSafePoint: 20}), + ) + got, err := w.RecvBatch(1) + require.NoError(t, err) + require.Equal(t, GCState{KeyspaceID: 7, TxnSafePoint: 10}, mustUpsert(t, got[0])) + + want := errors.New("watch terminated with pending initial and live changes") + w.cancel(want) + got, err = w.RecvBatch(3) + require.ErrorIs(t, err, want) + require.Nil(t, got) +} + +func queueLiveWhenInitialReceived(t *testing.T, w *GCStateWatcher, changes ...GCStateChange) { + t.Helper() + const name = "github.com/tikv/pd/pkg/gc/watchGCStatesInitialBatchReceived" + // Recreate the reachable merge state after select picks an initial batch + // while older live changes are queued. Enqueue in the hook only to force + // that branch deterministically, without depending on select randomness. + require.NoError(t, failpoint.EnableCall(name, func() { + for _, change := range changes { + w.liveCh <- change + } + changes = nil + })) + t.Cleanup(func() { require.NoError(t, failpoint.Disable(name)) }) +} + +func TestGCStateWatcherRemovedSuppressesInitial(t *testing.T) { + w := newGCStateWatcher(context.Background(), gcStateWatchConfig{initChannelCapacity: 1, liveChannelCapacity: 2}, false) + w.liveCh <- NewGCStateRemoved(7) + got, err := w.RecvBatch(1) + require.NoError(t, err) + removed, ok := got[0].RemovedKeyspaceID() + require.True(t, ok) + require.Equal(t, uint32(7), removed) + + w.initCh <- []GCStateChange{NewGCStateUpsert(GCState{KeyspaceID: 7})} + close(w.initCh) + _, ok, err = w.receiveOne(false) + require.NoError(t, err) + require.False(t, ok) + require.True(t, w.initDone) +} + +func TestGCStateWatcherDrainsBufferedInitBeforeReleasingDirtySet(t *testing.T) { + w := newGCStateWatcher(context.Background(), gcStateWatchConfig{initChannelCapacity: 1, liveChannelCapacity: 1}, false) + w.liveCh <- NewGCStateUpsert(GCState{KeyspaceID: 7}) + _, err := w.RecvBatch(1) + require.NoError(t, err) + w.initCh <- []GCStateChange{NewGCStateUpsert(GCState{KeyspaceID: 8})} + close(w.initCh) + + got, err := w.RecvBatch(1) + require.NoError(t, err) + require.Equal(t, uint32(8), mustUpsert(t, got[0]).KeyspaceID) + require.False(t, w.initDone) + require.NotNil(t, w.dirtyDuringInit) + + _, ok, err := w.receiveOne(false) + require.NoError(t, err) + require.False(t, ok) + require.True(t, w.initDone) + require.Nil(t, w.initCh) + require.Nil(t, w.dirtyDuringInit) +} + +func TestGCStateWatcherRecvBatchHonorsMaximum(t *testing.T) { + w := newGCStateWatcher(context.Background(), gcStateWatchConfig{liveChannelCapacity: 3}, true) + for id := uint32(1); id <= 3; id++ { + w.liveCh <- NewGCStateUpsert(GCState{KeyspaceID: id}) + } + got, err := w.RecvBatch(2) + require.NoError(t, err) + require.Len(t, got, 2) + got, err = w.RecvBatch(2) + require.NoError(t, err) + require.Len(t, got, 1) +} + +func TestGCStateWatcherCancellationDiscardsBufferedWork(t *testing.T) { + w := newGCStateWatcher(context.Background(), gcStateWatchConfig{liveChannelCapacity: 1}, true) + w.liveCh <- NewGCStateUpsert(GCState{KeyspaceID: 7}) + want := errors.New("watch terminated") + w.cancel(want) + got, err := w.RecvBatch(1) + require.ErrorIs(t, err, want) + require.Nil(t, got) +} + +func TestGCStateWatcherFirstCancellationCauseWins(t *testing.T) { + w := newGCStateWatcher(context.Background(), gcStateWatchConfig{liveChannelCapacity: 1}, true) + first := errors.New("first") + w.cancel(first) + w.cancel(errors.New("second")) + require.ErrorIs(t, w.Err(), first) +} + +func (s *gcStateManagerTestSuite) TestGCStateWatchPublishesAdvanceGCSafePoint() { + re := s.Require() + const keyspaceID = uint32(2) + _, err := s.manager.AdvanceTxnSafePoint(keyspaceID, 20, time.Now()) + re.NoError(err) + w, err := s.manager.WatchGCStates(context.Background(), true) + re.NoError(err) + defer w.Close() + + _, _, err = s.manager.AdvanceGCSafePoint(keyspaceID, 10) + re.NoError(err) + changes, err := w.RecvBatch(1) + re.NoError(err) + state := mustUpsert(s.T(), changes[0]) + re.Equal(GCState{KeyspaceID: keyspaceID, IsKeyspaceLevel: true, TxnSafePoint: 20, GCSafePoint: 10}, state) + re.Empty(state.GCBarriers) +} + +func (s *gcStateManagerTestSuite) TestGCStateWatchPublishesAdvanceTxnSafePoint() { + re := s.Require() + const keyspaceID = uint32(2) + w, err := s.manager.WatchGCStates(context.Background(), true) + re.NoError(err) + defer w.Close() + + _, err = s.manager.AdvanceTxnSafePoint(keyspaceID, 20, time.Now()) + re.NoError(err) + changes, err := w.RecvBatch(1) + re.NoError(err) + state := mustUpsert(s.T(), changes[0]) + re.Equal(GCState{KeyspaceID: keyspaceID, IsKeyspaceLevel: true, TxnSafePoint: 20}, state) + re.Empty(state.GCBarriers) +} + +func (s *gcStateManagerTestSuite) TestGCStateWatchCompatiblePathsPublishOnce() { + re := s.Require() + const keyspaceID = uint32(2) + _, err := s.manager.AdvanceTxnSafePoint(keyspaceID, 30, time.Now()) + re.NoError(err) + + gcWatcher, err := s.manager.WatchGCStates(context.Background(), true) + re.NoError(err) + _, _, err = s.manager.CompatibleUpdateGCSafePoint(keyspaceID, 10) + re.NoError(err) + changes, err := gcWatcher.RecvBatch(1) + re.NoError(err) + re.Len(changes, 1) + state := mustUpsert(s.T(), changes[0]) + re.Equal(GCState{KeyspaceID: keyspaceID, IsKeyspaceLevel: true, TxnSafePoint: 30, GCSafePoint: 10}, state) + re.Empty(state.GCBarriers) + re.Empty(gcWatcher.liveCh) + gcWatcher.Close() + + txnWatcher, err := s.manager.WatchGCStates(context.Background(), true) + re.NoError(err) + _, _, err = s.manager.CompatibleUpdateServiceGCSafePoint(keyspaceID, keypath.GCWorkerServiceSafePointID, 40, math.MaxInt64, time.Now()) + re.NoError(err) + changes, err = txnWatcher.RecvBatch(1) + re.NoError(err) + re.Len(changes, 1) + state = mustUpsert(s.T(), changes[0]) + re.Equal(GCState{KeyspaceID: keyspaceID, IsKeyspaceLevel: true, TxnSafePoint: 40, GCSafePoint: 10}, state) + re.Empty(state.GCBarriers) + re.Empty(txnWatcher.liveCh) + txnWatcher.Close() +} + +func (s *gcStateManagerTestSuite) TestGCStateWatchDoesNotPublishNoOpOrFailure() { + re := s.Require() + const keyspaceID = uint32(2) + _, err := s.manager.AdvanceTxnSafePoint(keyspaceID, 20, time.Now()) + re.NoError(err) + _, _, err = s.manager.AdvanceGCSafePoint(keyspaceID, 10) + re.NoError(err) + w, err := s.manager.WatchGCStates(context.Background(), true) + re.NoError(err) + defer w.Close() + + _, err = s.manager.AdvanceTxnSafePoint(keyspaceID, 20, time.Now()) + re.NoError(err) + _, _, err = s.manager.CompatibleUpdateGCSafePoint(keyspaceID, 10) + re.NoError(err) + _, _, err = s.manager.AdvanceGCSafePoint(keyspaceID, 9) + re.ErrorIs(err, errs.ErrDecreasingGCSafePoint) + re.Empty(w.liveCh) +} + +func (s *gcStateManagerTestSuite) TestGCStateWatchDoesNotPublishBarrierOnlyChanges() { + re := s.Require() + const keyspaceID = uint32(2) + _, err := s.manager.AdvanceTxnSafePoint(keyspaceID, 20, time.Now()) + re.NoError(err) + w, err := s.manager.WatchGCStates(context.Background(), true) + re.NoError(err) + defer w.Close() + + _, err = s.manager.SetGCBarrier(keyspaceID, "backup", 30, time.Hour, time.Now()) + re.NoError(err) + _, err = s.manager.DeleteGCBarrier(keyspaceID, "backup") + re.NoError(err) + re.Empty(w.liveCh) +} + +func (s *gcStateManagerTestSuite) TestGCStateWatchSlowConsumerIsolation() { + re := s.Require() + const keyspaceID = uint32(2) + watcherA, err := s.manager.registerGCStateWatcher(context.Background(), true, gcStateWatchConfig{liveChannelCapacity: 1}) + re.NoError(err) + defer watcherA.Close() + watcherB, err := s.manager.registerGCStateWatcher(context.Background(), true, gcStateWatchConfig{liveChannelCapacity: 4}) + re.NoError(err) + defer watcherB.Close() + + _, err = s.manager.AdvanceTxnSafePoint(keyspaceID, 10, time.Now()) + re.NoError(err) + changes, err := watcherB.RecvBatch(1) + re.NoError(err) + state := mustUpsert(s.T(), changes[0]) + re.Equal(GCState{KeyspaceID: keyspaceID, IsKeyspaceLevel: true, TxnSafePoint: 10}, state) + re.Empty(state.GCBarriers) + + _, err = s.manager.AdvanceTxnSafePoint(keyspaceID, 20, time.Now()) + re.NoError(err) + changes, err = watcherB.RecvBatch(1) + re.NoError(err) + state = mustUpsert(s.T(), changes[0]) + re.Equal(GCState{KeyspaceID: keyspaceID, IsKeyspaceLevel: true, TxnSafePoint: 20}, state) + re.Empty(state.GCBarriers) + + re.ErrorIs(watcherA.Err(), errs.ErrGCStateWatcherSlowConsumer) + re.NoError(watcherB.Err()) + re.NotContains(s.manager.watchers, watcherA.id) + re.Contains(s.manager.watchers, watcherB.id) + + reconnected, err := s.manager.WatchGCStates(context.Background(), false) + re.NoError(err) + defer reconnected.Close() + for { + changes, err = reconnected.RecvBatch(1) + re.NoError(err) + state = mustUpsert(s.T(), changes[0]) + if state.KeyspaceID == keyspaceID { + break + } + } + re.Equal(GCState{KeyspaceID: keyspaceID, IsKeyspaceLevel: true, TxnSafePoint: 20}, state) + re.Empty(state.GCBarriers) +} + +func (s *gcStateManagerTestSuite) TestGCStateWatcherMetrics() { + re := s.Require() + activeBefore := promtestutil.ToFloat64(gcStateWatcherGauge) + clientCancelBefore := promtestutil.ToFloat64(gcStateWatcherTerminationClientCancelCounter) + leaderLostBefore := promtestutil.ToFloat64(gcStateWatcherTerminationLeaderLostCounter) + slowConsumerBefore := promtestutil.ToFloat64(gcStateWatcherTerminationSlowConsumerCounter) + initErrorBefore := promtestutil.ToFloat64(gcStateWatcherTerminationInitErrorCounter) + assertTerminationDeltas := func(clientCancel, leaderLost, slowConsumer, initError float64) { + re.Equal(clientCancelBefore+clientCancel, promtestutil.ToFloat64(gcStateWatcherTerminationClientCancelCounter)) + re.Equal(leaderLostBefore+leaderLost, promtestutil.ToFloat64(gcStateWatcherTerminationLeaderLostCounter)) + re.Equal(slowConsumerBefore+slowConsumer, promtestutil.ToFloat64(gcStateWatcherTerminationSlowConsumerCounter)) + re.Equal(initErrorBefore+initError, promtestutil.ToFloat64(gcStateWatcherTerminationInitErrorCounter)) + } + + stop := s.manager.OnNodeBecomesLeader() + w, err := s.manager.WatchGCStates(context.Background(), true) + re.NoError(err) + re.Equal(activeBefore+1, promtestutil.ToFloat64(gcStateWatcherGauge)) + + stop() + re.ErrorIs(w.Err(), errs.ErrNotLeader) + re.Equal(activeBefore, promtestutil.ToFloat64(gcStateWatcherGauge)) + assertTerminationDeltas(0, 1, 0, 0) + w.Close() + assertTerminationDeltas(0, 1, 0, 0) + + stopRemainingCases := s.manager.OnNodeBecomesLeader() + defer stopRemainingCases() + + clientCanceled, err := s.manager.WatchGCStates(context.Background(), true) + re.NoError(err) + re.Equal(activeBefore+1, promtestutil.ToFloat64(gcStateWatcherGauge)) + clientCanceled.Close() + clientCanceled.Close() + re.Equal(activeBefore, promtestutil.ToFloat64(gcStateWatcherGauge)) + assertTerminationDeltas(1, 1, 0, 0) + + slowConsumer, err := s.manager.registerGCStateWatcher(context.Background(), true, gcStateWatchConfig{liveChannelCapacity: 1}) + re.NoError(err) + re.Equal(activeBefore+1, promtestutil.ToFloat64(gcStateWatcherGauge)) + _, err = s.manager.AdvanceTxnSafePoint(2, 10, time.Now()) + re.NoError(err) + re.Equal(activeBefore+1, promtestutil.ToFloat64(gcStateWatcherGauge)) + _, err = s.manager.AdvanceTxnSafePoint(2, 20, time.Now()) + re.NoError(err) + re.ErrorIs(slowConsumer.Err(), errs.ErrGCStateWatcherSlowConsumer) + re.Equal(activeBefore, promtestutil.ToFloat64(gcStateWatcherGauge)) + assertTerminationDeltas(1, 1, 1, 0) + slowConsumer.Close() + assertTerminationDeltas(1, 1, 1, 0) + + const errorMessage = "injected initial watch failure" + func() { + re.NoError(failpoint.Enable("github.com/tikv/pd/pkg/gc/iterateAllKeyspacesGCStatesError", fmt.Sprintf(`return(%q)`, errorMessage))) + defer func() { re.NoError(failpoint.Disable("github.com/tikv/pd/pkg/gc/iterateAllKeyspacesGCStatesError")) }() + initFailed, err := s.manager.WatchGCStates(context.Background(), false) + re.NoError(err) + _, err = initFailed.RecvBatch(1) + re.ErrorContains(err, errorMessage) + re.Equal(activeBefore, promtestutil.ToFloat64(gcStateWatcherGauge)) + assertTerminationDeltas(1, 1, 1, 1) + initFailed.Close() + assertTerminationDeltas(1, 1, 1, 1) + }() +} + +func mustUpsert(t testing.TB, change GCStateChange) GCState { + state, ok := change.Upsert() + require.True(t, ok) + return state +} + +func TestGCStateWatcherDonePublishesFirstCause(t *testing.T) { + ctx, cancel := context.WithCancelCause(context.Background()) + defer cancel(context.Canceled) + w := newGCStateWatcher(ctx, gcStateWatchConfig{initChannelCapacity: 1, liveChannelCapacity: 1}, true) + defer w.Close() + done := w.Done() + require.Equal(t, done, w.Done()) + select { + case <-done: + require.FailNow(t, "watcher terminated before cancellation") + default: + } + cancel(errs.ErrNotLeader) + select { + case <-done: + case <-time.After(5 * time.Second): + require.FailNow(t, "watcher termination was not notified") + } + require.ErrorIs(t, w.Err(), errs.ErrNotLeader) + w.Close() + require.ErrorIs(t, w.Err(), errs.ErrNotLeader) + require.Equal(t, done, w.Done()) +} + +func TestGCStateWatcherUsesEnabledMetadataCache(t *testing.T) { + _, _, manager, clean, cancel := newGCStateManagerForTest(t, newGCStateManagerForTestOptions{useEnabledKeyspaceCache: true}) + defer clean() + defer cancel() + ctx, stop := context.WithTimeout(context.Background(), 10*time.Second) + defer stop() + require.NoError(t, manager.enabledKeyspaces.waitReady(ctx)) + + id := uint32(19) + _, err := manager.keyspaceManager.CreateKeyspaceByID(&keyspace.CreateKeyspaceByIDRequest{ + ID: &id, Name: "watch-cache-created", Config: map[string]string{keyspace.GCManagementType: keyspace.KeyspaceLevelGC}, CreateTime: time.Now().Unix(), + }) + require.NoError(t, err) + // A full watcher must use the index even when the legacy iterator fails. + const failpointName = "github.com/tikv/pd/pkg/gc/iterateAllKeyspacesGCStatesError" + require.NoError(t, failpoint.Enable(failpointName, `return("legacy iterator used")`)) + defer func() { require.NoError(t, failpoint.Disable(failpointName)) }() + w, err := manager.WatchGCStates(ctx, false) + require.NoError(t, err) + defer w.Close() + for { + changes, err := w.RecvBatch(16) + require.NoError(t, err) + for _, change := range changes { + if state := mustUpsert(t, change); state.KeyspaceID == id { + require.True(t, state.IsKeyspaceLevel) + return + } + } + } +} + +func TestGCStateWatcherWaitsForCommittedMetadataRevision(t *testing.T) { + _, _, manager, clean, cancel := newGCStateManagerForTest(t, newGCStateManagerForTestOptions{useEnabledKeyspaceCache: true}) + defer clean() + defer cancel() + ctx, stop := context.WithTimeout(context.Background(), 10*time.Second) + defer stop() + + // Replace the test term's cache with one whose real etcd watch cannot + // start until this test releases it. Its initial snapshot is ready, but + // the subsequent metadata commit remains unapplied at registration. + watchStarted := make(chan struct{}, 16) + releaseWatch := make(chan struct{}) + var releaseOnce sync.Once + release := func() { releaseOnce.Do(func() { close(releaseWatch) }) } + defer release() + manager.mu.Lock() + manager.cancelEnabledKeyspaces() + termCtx, termCancel := context.WithCancel(context.Background()) + cache := newEnabledKeyspaceCache(termCtx, manager.etcdClient, keypath.KeyspaceMetaPrefix()) + cache.watcherFactory = func(client *clientv3.Client) clientv3.Watcher { + return &pauseBeforeWatch{ + Watcher: clientv3.NewWatcher(client), + started: watchStarted, + release: releaseWatch, + } + } + manager.enabledKeyspaces = cache + manager.cancelEnabledKeyspaces = termCancel + manager.mu.Unlock() + cacheDone := make(chan struct{}) + go func() { + cache.run() + close(cacheDone) + }() + defer func() { + termCancel() + select { + case <-cacheDone: + case <-time.After(5 * time.Second): + t.Error("replacement cache did not stop") + } + }() + require.NoError(t, cache.waitReady(ctx)) + select { + case <-watchStarted: + case <-ctx.Done(): + t.Fatal("replacement cache did not reach the watch") + } + + id := uint32(19) + _, err := manager.keyspaceManager.CreateKeyspaceByID(&keyspace.CreateKeyspaceByIDRequest{ + ID: &id, Name: "watch-cache-lagged", Config: map[string]string{keyspace.GCManagementType: keyspace.KeyspaceLevelGC}, CreateTime: time.Now().Unix(), + }) + require.NoError(t, err) + commit, err := manager.etcdClient.Get(ctx, keypath.KeyspaceMetaPath(id)) + require.NoError(t, err) + require.Less(t, cache.appliedRevision(), commit.Header.Revision) + + probed := make(chan int64, 1) + const hook = "github.com/tikv/pd/pkg/gc/watchGCStatesTargetRevisionProbed" + require.NoError(t, failpoint.EnableCall(hook, func(revision int64) { probed <- revision })) + defer func() { require.NoError(t, failpoint.Disable(hook)) }() + w, err := manager.WatchGCStates(ctx, false) + require.NoError(t, err) + defer w.Close() + select { + case target := <-probed: + require.GreaterOrEqual(t, target, commit.Header.Revision) + case <-ctx.Done(): + t.Fatal("watcher did not probe the target revision") + } + received := make(chan error, 1) + go func() { + _, err := w.RecvBatch(1) + received <- err + }() + select { + case err := <-received: + t.Fatalf("initial state arrived before the metadata watch resumed: %v", err) + case <-time.After(200 * time.Millisecond): + } + + release() + require.NoError(t, <-received) + for { + changes, err := w.RecvBatch(16) + require.NoError(t, err) + for _, change := range changes { + if state := mustUpsert(t, change); state.KeyspaceID == id { + require.True(t, state.IsKeyspaceLevel) + return + } + } + } +} + +func TestGCStateWatcherWaitReadyStopsOnCancelAndLeaderLoss(t *testing.T) { + _, _, manager, clean, cancel := newGCStateManagerForTest(t, newGCStateManagerForTestOptions{ + useEnabledKeyspaceCache: true, + beforeLeader: func(c *clientv3.Client) { + _, err := c.Put(context.Background(), keypath.KeyspaceMetaPath(19), "invalid protobuf") + require.NoError(t, err) + }, + }) + defer clean() + defer cancel() + + ctx, stop := context.WithCancel(context.Background()) + result := make(chan error, 1) + go func() { + _, err := manager.WatchGCStates(ctx, false) + result <- err + }() + stop() + require.ErrorIs(t, <-result, context.Canceled) + require.Empty(t, manager.watchers) + + waiting := make(chan struct{}, 1) + const waitHook = "github.com/tikv/pd/pkg/gc/watchGCStatesBeforeCacheReady" + require.NoError(t, failpoint.EnableCall(waitHook, func() { waiting <- struct{}{} })) + defer func() { require.NoError(t, failpoint.Disable(waitHook)) }() + result = make(chan error, 1) + go func() { + _, err := manager.WatchGCStates(context.Background(), false) + result <- err + }() + select { + case <-waiting: + case <-time.After(5 * time.Second): + t.Fatal("watcher did not begin waiting for cache readiness") + } + stopTerm := manager.OnNodeBecomesLeader() + defer stopTerm() + select { + case err := <-result: + require.ErrorIs(t, err, errs.ErrNotLeader) + case <-time.After(5 * time.Second): + t.Fatal("watcher did not exit on leader change") + } + require.Empty(t, manager.watchers) +} + +func TestGCStateWatcherIndexedInitialMergesConcurrentGCWrite(t *testing.T) { + _, _, manager, clean, cancel := newGCStateManagerForTest(t, newGCStateManagerForTestOptions{useEnabledKeyspaceCache: true}) + defer clean() + defer cancel() + ctx, stop := context.WithTimeout(context.Background(), 10*time.Second) + defer stop() + require.NoError(t, manager.enabledKeyspaces.waitReady(ctx)) + _, err := manager.AdvanceTxnSafePoint(2, 10, time.Now()) + require.NoError(t, err) + + reached := make(chan struct{}) + release := make(chan struct{}) + var once sync.Once + releaseLoader := func() { once.Do(func() { close(release) }) } + defer releaseLoader() + const hook = "github.com/tikv/pd/pkg/gc/watchGCStatesInitialStateLoaded" + require.NoError(t, failpoint.EnableCall(hook, func(id uint32) { + if id == 2 { + close(reached) + <-release + } + })) + defer func() { require.NoError(t, failpoint.Disable(hook)) }() + w, err := manager.WatchGCStates(ctx, false) + require.NoError(t, err) + defer w.Close() + select { + case <-reached: + case <-ctx.Done(): + t.Fatal("initial loader did not reach keyspace 2") + } + _, err = manager.AdvanceTxnSafePoint(2, 20, time.Now()) + require.NoError(t, err) + for { + changes, err := w.RecvBatch(1) + require.NoError(t, err) + if state := mustUpsert(t, changes[0]); state.KeyspaceID == 2 { + require.Equal(t, uint64(20), state.TxnSafePoint) + break + } + } + releaseLoader() + require.Eventually(t, func() bool { + for { + change, ok, err := w.receiveOne(false) + require.NoError(t, err) + if !ok { + return w.initDone + } + state := mustUpsert(t, change) + require.False(t, state.KeyspaceID == 2 && state.TxnSafePoint == 10) + } + }, 5*time.Second, 10*time.Millisecond) +} + +func TestGCStateWatcherIndexedPostRegistrationErrorCleansUp(t *testing.T) { + _, _, manager, clean, cancel := newGCStateManagerForTest(t, newGCStateManagerForTestOptions{useEnabledKeyspaceCache: true}) + defer clean() + defer cancel() + ctx, stop := context.WithTimeout(context.Background(), 10*time.Second) + defer stop() + require.NoError(t, manager.enabledKeyspaces.waitReady(ctx)) + const hook = "github.com/tikv/pd/pkg/gc/watchGCStatesRegistered" + require.NoError(t, failpoint.EnableCall(hook, manager.cancelEnabledKeyspaces)) + defer func() { require.NoError(t, failpoint.Disable(hook)) }() + w, err := manager.WatchGCStates(ctx, false) + require.NoError(t, err) + defer w.Close() + _, err = w.RecvBatch(1) + require.Error(t, err) + require.NotContains(t, manager.watchers, w.id) +} diff --git a/pkg/gc/metrics.go b/pkg/gc/metrics.go index 86e4f088150..912d99a9575 100644 --- a/pkg/gc/metrics.go +++ b/pkg/gc/metrics.go @@ -78,12 +78,32 @@ var ( gcStateCacheAccessHitCounter = gcStateCacheAccessCounter.WithLabelValues("hit") gcStateCacheAccessSlowHitCounter = gcStateCacheAccessCounter.WithLabelValues("slow_hit") gcStateCacheAccessMissCounter = gcStateCacheAccessCounter.WithLabelValues("miss") + + gcStateWatcherGauge = prometheus.NewGauge(prometheus.GaugeOpts{ + Namespace: "pd", + Subsystem: "gc", + Name: "watcher_count", + Help: "Current number of active GC state watchers.", + }) + gcStateWatcherTerminationCounter = prometheus.NewCounterVec(prometheus.CounterOpts{ + Namespace: "pd", + Subsystem: "gc", + Name: "watcher_termination_total", + Help: "Total number of GC state watcher terminations by reason.", + }, []string{"reason"}) + + gcStateWatcherTerminationClientCancelCounter = gcStateWatcherTerminationCounter.WithLabelValues("client_cancel") + gcStateWatcherTerminationLeaderLostCounter = gcStateWatcherTerminationCounter.WithLabelValues("leader_lost") + gcStateWatcherTerminationSlowConsumerCounter = gcStateWatcherTerminationCounter.WithLabelValues("slow_consumer") + gcStateWatcherTerminationInitErrorCounter = gcStateWatcherTerminationCounter.WithLabelValues("init_error") ) func init() { prometheus.MustRegister(productionBarrierMetrics) prometheus.MustRegister(gcSafePointGauge) prometheus.MustRegister(gcStateCacheAccessCounter) + prometheus.MustRegister(gcStateWatcherGauge) + prometheus.MustRegister(gcStateWatcherTerminationCounter) } type barrierMetricScope struct { diff --git a/pkg/gc/metrics_test.go b/pkg/gc/metrics_test.go index b218ab47d7b..012cb6d7bf0 100644 --- a/pkg/gc/metrics_test.go +++ b/pkg/gc/metrics_test.go @@ -286,22 +286,23 @@ func TestBarrierMetricsRegistrationAndLeadership(t *testing.T) { observe := func(m *GCStateManager, id uint32) []barrierWarning { return m.barrierMetrics.observeMetrics(m.barrierMetrics.generation(), id, "tenant", barriers, nil, now) } - first.OnNodeBecomesLeader() + stopFirst := first.OnNodeBecomesLeader() require.Len(t, observe(first, 42), 1) require.Contains(t, gatherBarrierMetrics(t, prometheus.DefaultGatherer), "keyspace/42/old") generation := first.barrierMetrics.generation() - first.OnNodeBecomesLeader() + stopReplacement := first.OnNodeBecomesLeader() + defer stopReplacement() require.Empty(t, gatherBarrierMetrics(t, prometheus.DefaultGatherer)) require.Empty(t, first.barrierMetrics.observeMetrics(generation, 42, "tenant-a", barriers, nil, now)) require.Empty(t, gatherBarrierMetrics(t, prometheus.DefaultGatherer), "stale pre-leadership read cannot publish") - first.OnNodeBecomesFollower() // Previous lease ends, one leadership remains. + stopFirst() // Previous lease ends, the replacement leadership remains. require.Len(t, observe(first, 42), 1) require.Contains(t, gatherBarrierMetrics(t, prometheus.DefaultGatherer), "keyspace/42/old") - second.OnNodeBecomesLeader() + stopSecond := second.OnNodeBecomesLeader() require.Len(t, observe(second, 43), 1) first.CloseBarrierMetrics() require.Equal(t, map[string]float64{"keyspace/43/old": 1_999_712_000}, gatherBarrierMetrics(t, prometheus.DefaultGatherer), "old owner cleanup must preserve replacement") - second.OnNodeBecomesFollower() + stopSecond() require.Empty(t, gatherBarrierMetrics(t, prometheus.DefaultGatherer)) require.Nil(t, productionBarrierMetrics.current.Load(), "registry must not retain a stopped manager") } @@ -811,6 +812,7 @@ func (s *gcStateManagerTestSuite) TestBarrierMetricsRemovalFencesInflightPublica re := s.Require() now := time.Unix(2_000_000_000, 0) m := s.manager + stop := m.OnNodeBecomesLeader() m.barrierMetrics.now = func() time.Time { return now } registry := prometheus.NewRegistry() registry.MustRegister(m.barrierMetrics) @@ -829,7 +831,7 @@ func (s *gcStateManagerTestSuite) TestBarrierMetricsRemovalFencesInflightPublica advance() re.Equal(expected, gatherBarrierMetrics(s.T(), registry), "every accepted metadata state remains observable") } - m.OnNodeBecomesFollower() + stop() re.Empty(gatherBarrierMetrics(s.T(), registry)) m.OnNodeBecomesLeader() re.Empty(gatherBarrierMetrics(s.T(), registry)) diff --git a/pkg/storage/endpoint/gc_states.go b/pkg/storage/endpoint/gc_states.go index 439bb0ca914..00ae63b56d8 100644 --- a/pkg/storage/endpoint/gc_states.go +++ b/pkg/storage/endpoint/gc_states.go @@ -472,6 +472,16 @@ func (p GCStateProvider) RunInGCStateTransaction(f func(wb *GCStateWriteBatch) e OpType: kv.RawTxnOpPut, Value: nextRevision, }) + } else { + // etcd treats a Compare-only transaction as serializable and may check + // the revision on a stale follower. A non-serializable Get makes etcd + // linearize the whole transaction before evaluating the comparison, + // including when it takes the empty Else branch. It does not write or + // advance the revision; its response is counted along with the ops below. + ops = append(ops, kv.RawTxnOp{ + Key: revisionKey, + OpType: kv.RawTxnOpGet, + }) } txn, err := p.storage.createRawTxn() diff --git a/pkg/storage/endpoint/gc_states_txn_test.go b/pkg/storage/endpoint/gc_states_txn_test.go new file mode 100644 index 00000000000..733f7e0b940 --- /dev/null +++ b/pkg/storage/endpoint/gc_states_txn_test.go @@ -0,0 +1,181 @@ +// Copyright 2026 TiKV Project Authors. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package endpoint + +import ( + "context" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" + clientv3 "go.etcd.io/etcd/client/v3" + + "github.com/pingcap/errors" + + "github.com/tikv/pd/pkg/errs" + "github.com/tikv/pd/pkg/storage/kv" + "github.com/tikv/pd/pkg/utils/etcdutil" + "github.com/tikv/pd/pkg/utils/keypath" +) + +func TestGCStateReadOnlyTransactionLaggingFollower(t *testing.T) { + for _, concurrentWrite := range []bool{false, true} { + name := "completed-write" + if concurrentWrite { + name = "concurrent-write" + } + t.Run(name, func(t *testing.T) { + re := require.New(t) + servers, _, clean := etcdutil.NewTestEtcdCluster(t, 3, nil) + defer clean() + ctx, cancel := context.WithTimeout(context.Background(), 20*time.Second) + defer cancel() + + leader, follower := servers[0], servers[1] + for _, server := range servers { + if uint64(server.Server.ID()) == server.Server.Lead() { + leader = server + } else { + follower = server + } + } + re.NotEqual(leader.Server.ID(), follower.Server.ID()) + newClient := func(endpoint string) *clientv3.Client { + client, err := clientv3.New(clientv3.Config{Endpoints: []string{endpoint}, Context: ctx}) + re.NoError(err) + t.Cleanup(func() { re.NoError(client.Close()) }) + return client + } + freshClient := newClient(leader.Config().ListenClientUrls[0].String()) + staleClient := newClient(follower.Config().ListenClientUrls[0].String()) + freshKV, staleKV := kv.NewEtcdKVBase(freshClient), kv.NewEtcdKVBase(staleClient) + writer := NewStorageEndpoint(freshKV, nil).GetGCStateProvider() + // Pin ordinary reads to the leader and validation to the follower. This + // reproduces the problematic round-robin routing without relying on chance. + reader := NewStorageEndpoint(struct { + kv.Base + kv.RawTxnCapable + }{freshKV, staleKV}, nil).GetGCStateProvider() + write := func() error { + return writer.RunInGCStateTransaction(func(wb *GCStateWriteBatch) error { + return wb.SetGCSafePoint(0, 100) + }) + } + re.NoError(write()) + revisionKey := keypath.GCStateRevisionPath() + before, err := staleClient.Get(ctx, revisionKey) + re.NoError(err) + re.Len(before.Kvs, 1) + re.Equal("1", string(before.Kvs[0].Value)) + + // Isolate only the follower's Raft traffic; client RPCs still work and + // the other two members can commit the next (and final) write. + for _, server := range servers { + if server != follower { + server.Server.CutPeer(follower.Server.ID()) + follower.Server.CutPeer(server.Server.ID()) + } + } + // MendPeer restarts remote pipelines and must only be called once. + mend := sync.OnceFunc(func() { + for _, server := range servers { + if server != follower { + server.Server.MendPeer(follower.Server.ID()) + follower.Server.MendPeer(server.Server.ID()) + } + } + }) + defer mend() + if !concurrentWrite { + re.NoError(write()) + } + + ready := make(chan struct{}) + done := make(chan error, 1) + go func() { + done <- reader.RunInGCStateTransaction(func(_ *GCStateWriteBatch) error { + defer close(ready) + if concurrentWrite { + // Commit after the reader has sampled revision 1, so validation + // must reject it even though the follower still has revision 1. + return write() + } + return nil + }) + }() + select { + case <-ready: + case <-ctx.Done(): + t.Fatal(ctx.Err()) + } + fresh, err := freshClient.Get(ctx, revisionKey) + re.NoError(err) + re.Len(fresh.Kvs, 1) + re.Equal("2", string(fresh.Kvs[0].Value)) + stale, err := staleClient.Get(ctx, revisionKey, clientv3.WithSerializable()) + re.NoError(err) + re.Equal(before.Kvs, stale.Kvs, "the follower must still be behind after all writes finish") + + // With no new writes, validation must wait for the follower to catch up. + // An empty Compare-only transaction instead returns immediately using + // revision 1: a false conflict, or a missed real conflict, respectively. + var txnErr error + select { + case txnErr = <-done: + mend() + case <-time.After(500 * time.Millisecond): + mend() + select { + case txnErr = <-done: + case <-ctx.Done(): + t.Fatal(ctx.Err()) + } + } + if concurrentWrite { + re.True(errors.ErrorEqual(txnErr, errs.ErrEtcdTxnConflict), "got %v", txnErr) + } else { + re.NoError(txnErr) + } + after, err := freshClient.Get(ctx, revisionKey) + re.NoError(err) + re.Equal(fresh.Kvs, after.Kvs, "read-only validation must not write the GC revision") + re.Equal(fresh.Header.Revision, after.Header.Revision, "read-only validation must not advance etcd revision") + }) + } +} + +func TestGCStateReadOnlyTransactionRevision(t *testing.T) { + _, client, clean := etcdutil.NewTestEtcdCluster(t, 1, nil) + defer clean() + re := require.New(t) + provider := NewStorageEndpoint(kv.NewEtcdKVBase(client), nil).GetGCStateProvider() + for _, initialized := range []bool{false, true} { + if initialized { + re.NoError(provider.RunInGCStateTransaction(func(wb *GCStateWriteBatch) error { + return wb.SetGCSafePoint(0, 100) + })) + } + before, err := client.Get(client.Ctx(), keypath.GCStateRevisionPath()) + re.NoError(err) + for range 2 { + re.NoError(provider.RunInGCStateTransaction(func(_ *GCStateWriteBatch) error { return nil })) + } + after, err := client.Get(client.Ctx(), keypath.GCStateRevisionPath()) + re.NoError(err) + re.Equal(before.Kvs, after.Kvs) + re.Equal(before.Header.Revision, after.Header.Revision) + } +} diff --git a/server/cluster/cluster.go b/server/cluster/cluster.go index a37c31a2c60..0192faabd62 100644 --- a/server/cluster/cluster.go +++ b/server/cluster/cluster.go @@ -492,8 +492,7 @@ func (c *RaftCluster) Start(s Server, bootstrap bool) (err error) { go c.startProgressGC() go c.runStorageSizeCollector(s.GetMeteringWriter(), c.regionLabeler, s.GetKeyspaceManager()) - s.GetGCStateManager().OnNodeBecomesLeader() - c.stopGCStateManager = s.GetGCStateManager().OnNodeBecomesFollower + c.stopGCStateManager = s.GetGCStateManager().OnNodeBecomesLeader() log.Info("start background jobs completed", zap.Duration("cost", time.Since(backgroundJobsStart))) runnersStart := time.Now() diff --git a/server/gc_service.go b/server/gc_service.go index 122ddc9d5d2..f4df26e59ba 100644 --- a/server/gc_service.go +++ b/server/gc_service.go @@ -16,6 +16,7 @@ package server import ( "context" + "errors" "math" "time" @@ -28,14 +29,27 @@ import ( "github.com/pingcap/kvproto/pkg/pdpb" "github.com/pingcap/log" + "github.com/tikv/pd/pkg/errs" "github.com/tikv/pd/pkg/gc" "github.com/tikv/pd/pkg/keyspace/constant" "github.com/tikv/pd/pkg/storage/endpoint" "github.com/tikv/pd/pkg/utils/grpcutil" + "github.com/tikv/pd/pkg/utils/logutil" "github.com/tikv/pd/pkg/utils/tsoutil" "github.com/tikv/pd/pkg/utils/typeutil" ) +const ( + watchGCStatesRecvBatchSize = 1024 + maxWatchGCStatesResponseSize = 1 << 20 +) + +type gcStateChangeReceiver interface { + RecvBatch(maxChanges int) ([]gc.GCStateChange, error) + Done() <-chan struct{} + Err() error +} + // UpdateGCSafePoint implements gRPC PDServer. // // Deprecated: Use AdvanceGCSafePoint instead. Note that it's only for use of GC internal. @@ -533,6 +547,118 @@ func gcStateToProto(gcState gc.GCState, now time.Time) *pdpb.GCState { } } +func gcStateChangeToProto(change gc.GCStateChange) (*pdpb.GCStateChange, error) { + if state, ok := change.Upsert(); ok { + state.GCBarriers = nil + return &pdpb.GCStateChange{Change: &pdpb.GCStateChange_Upsert{ + Upsert: gcStateToProto(state, time.Time{}), + }}, nil + } + if keyspaceID, ok := change.RemovedKeyspaceID(); ok { + return &pdpb.GCStateChange{Change: &pdpb.GCStateChange_Removed{ + Removed: &pdpb.KeyspaceScope{ + Keyspace: &pdpb.KeyspaceScope_KeyspaceId{KeyspaceId: keyspaceID}, + }, + }}, nil + } + return nil, errors.New("invalid GC state change") +} + +func splitWatchGCStatesResponses(changes []*pdpb.GCStateChange, maxSize int) []*pdpb.WatchGCStatesResponse { + if len(changes) == 0 { + return nil + } + + responses := make([]*pdpb.WatchGCStatesResponse, 0, 1) + newResponse := func() (*pdpb.WatchGCStatesResponse, int) { + response := &pdpb.WatchGCStatesResponse{Header: grpcutil.WrapHeader()} + return response, response.Size() + } + current, currentSize := newResponse() + for _, change := range changes { + changeSize := (&pdpb.WatchGCStatesResponse{Changes: []*pdpb.GCStateChange{change}}).Size() + if len(current.Changes) > 0 && currentSize+changeSize > maxSize { + responses = append(responses, current) + current, currentSize = newResponse() + } + + current.Changes = append(current.Changes, change) + currentSize += changeSize + if len(current.Changes) == 1 && currentSize > maxSize { + log.Warn("GC state change exceeds WatchGCStates response size", + zap.Int("serialized-size", currentSize), + zap.Int("max-size", maxSize)) + responses = append(responses, current) + current, currentSize = newResponse() + } + } + if len(current.Changes) > 0 { + responses = append(responses, current) + } + return responses +} + +func watchGCStatesErrorToStatus(err error) error { + switch { + case errors.Is(err, errs.ErrGCStateWatcherSlowConsumer): + return status.Error(codes.ResourceExhausted, err.Error()) + case errors.Is(err, errs.ErrNotLeader): + return errs.ErrNotLeader + case errors.Is(err, context.Canceled), errors.Is(err, context.DeadlineExceeded): + return status.FromContextError(err).Err() + default: + return status.Error(codes.Unavailable, err.Error()) + } +} + +func serveWatchGCStates(receiver gcStateChangeReceiver, stream pdpb.PD_WatchGCStatesServer, maxResponseSize int) error { + if err := receiver.Err(); err != nil { + return watchGCStatesErrorToStatus(err) + } + resultCh := make(chan error, 1) + go func() { + defer logutil.LogPanic() + resultCh <- sendWatchGCStates(receiver, stream, maxResponseSize) + }() + // Do not join the worker here: returning lets gRPC tear down the transport + // stream, which interrupts a Send blocked on flow control. + select { + case err := <-resultCh: + if cause := receiver.Err(); cause != nil { + return watchGCStatesErrorToStatus(cause) + } + return err + case <-receiver.Done(): + return watchGCStatesErrorToStatus(receiver.Err()) + } +} + +func sendWatchGCStates(receiver gcStateChangeReceiver, stream pdpb.PD_WatchGCStatesServer, maxResponseSize int) error { + for { + changes, err := receiver.RecvBatch(watchGCStatesRecvBatchSize) + if err != nil { + return watchGCStatesErrorToStatus(err) + } + protoChanges := make([]*pdpb.GCStateChange, 0, len(changes)) + for _, change := range changes { + converted, err := gcStateChangeToProto(change) + if err != nil { + log.Error("failed to convert GC state change", zap.Error(err)) + return status.Error(codes.Internal, err.Error()) + } + protoChanges = append(protoChanges, converted) + } + for _, response := range splitWatchGCStatesResponses(protoChanges, maxResponseSize) { + if err := receiver.Err(); err != nil { + return watchGCStatesErrorToStatus(err) + } + if err := stream.Send(response); err != nil { + return err + } + } + } +} + // AdvanceGCSafePoint tries to advance the GC safe point. func (s *GrpcServer) AdvanceGCSafePoint(ctx context.Context, request *pdpb.AdvanceGCSafePointRequest) (*pdpb.AdvanceGCSafePointResponse, error) { done, err := s.rateLimitCheck() @@ -806,6 +932,30 @@ func (s *GrpcServer) GetAllKeyspacesGCStates(ctx context.Context, request *pdpb. }, nil } +// WatchGCStates streams effective GC state changes from this PD server. +func (s *GrpcServer) WatchGCStates(request *pdpb.WatchGCStatesRequest, stream pdpb.PD_WatchGCStatesServer) error { + done, err := s.rateLimitCheck() + if err != nil { + return err + } + if done != nil { + defer done() + } + if err := s.validateRequest(request.GetHeader()); err != nil { + return err + } + if s.GetRaftCluster() == nil { + return status.Error(codes.Unavailable, errs.ErrNotBootstrapped.FastGenByArgs().Error()) + } + + watcher, err := s.gcStateManager.WatchGCStates(stream.Context(), request.GetSkipLoadingInitial()) + if err != nil { + return watchGCStatesErrorToStatus(err) + } + defer watcher.Close() + return serveWatchGCStates(watcher, stream, maxWatchGCStatesResponseSize) +} + // SetGlobalGCBarrier sets a global GC barrier. func (s *GrpcServer) SetGlobalGCBarrier(ctx context.Context, request *pdpb.SetGlobalGCBarrierRequest) (*pdpb.SetGlobalGCBarrierResponse, error) { done, err := s.rateLimitCheck() diff --git a/server/gc_service_test.go b/server/gc_service_test.go new file mode 100644 index 00000000000..50f09f98113 --- /dev/null +++ b/server/gc_service_test.go @@ -0,0 +1,498 @@ +// Copyright 2026 TiKV Project Authors. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package server + +import ( + "context" + "errors" + "math" + "net" + "testing" + "time" + + "github.com/stretchr/testify/require" + "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/credentials/insecure" + "google.golang.org/grpc/metadata" + "google.golang.org/grpc/status" + "google.golang.org/grpc/test/bufconn" + + "github.com/pingcap/kvproto/pkg/pdpb" + + "github.com/tikv/pd/pkg/errs" + "github.com/tikv/pd/pkg/gc" + "github.com/tikv/pd/pkg/storage/endpoint" + "github.com/tikv/pd/pkg/utils/grpcutil" +) + +func TestGCStateChangeToProto(t *testing.T) { + testCases := []struct { + name string + change gc.GCStateChange + want *pdpb.GCStateChange + wantErr bool + }{ + { + name: "complete upsert", + change: gc.NewGCStateUpsert(gc.GCState{ + KeyspaceID: 7, + IsKeyspaceLevel: true, + TxnSafePoint: 10, + GCSafePoint: 5, + GCBarriers: []*endpoint.GCBarrier{ + {BarrierID: "test-barrier", BarrierTS: 8}, + }, + }), + want: &pdpb.GCStateChange{Change: &pdpb.GCStateChange_Upsert{Upsert: &pdpb.GCState{ + KeyspaceScope: &pdpb.KeyspaceScope{Keyspace: &pdpb.KeyspaceScope_KeyspaceId{KeyspaceId: 7}}, + IsKeyspaceLevelGc: true, + TxnSafePoint: 10, + GcSafePoint: 5, + }}}, + }, + { + name: "removed scope", + change: gc.NewGCStateRemoved(9), + want: &pdpb.GCStateChange{Change: &pdpb.GCStateChange_Removed{Removed: &pdpb.KeyspaceScope{ + Keyspace: &pdpb.KeyspaceScope_KeyspaceId{KeyspaceId: 9}, + }}}, + }, + { + name: "invalid zero value", + wantErr: true, + }, + } + + for _, testCase := range testCases { + t.Run(testCase.name, func(t *testing.T) { + got, err := gcStateChangeToProto(testCase.change) + if testCase.wantErr { + require.Error(t, err) + require.Nil(t, got) + return + } + require.NoError(t, err) + switch want := testCase.want.GetChange().(type) { + case *pdpb.GCStateChange_Upsert: + upsert := got.GetUpsert() + require.NotNil(t, upsert) + require.Equal(t, want.Upsert.GetKeyspaceScope().GetKeyspaceId(), upsert.GetKeyspaceScope().GetKeyspaceId()) + require.Equal(t, want.Upsert.GetIsKeyspaceLevelGc(), upsert.GetIsKeyspaceLevelGc()) + require.Equal(t, want.Upsert.GetTxnSafePoint(), upsert.GetTxnSafePoint()) + require.Equal(t, want.Upsert.GetGcSafePoint(), upsert.GetGcSafePoint()) + require.Empty(t, upsert.GetGcBarriers()) + require.Nil(t, got.GetRemoved()) + case *pdpb.GCStateChange_Removed: + removed := got.GetRemoved() + require.NotNil(t, removed) + require.Equal(t, want.Removed.GetKeyspaceId(), removed.GetKeyspaceId()) + require.Nil(t, got.GetUpsert()) + default: + require.FailNow(t, "unexpected expected GC state change type") + } + }) + } +} + +func TestSplitWatchGCStatesResponses(t *testing.T) { + change := &pdpb.GCStateChange{Change: &pdpb.GCStateChange_Upsert{Upsert: &pdpb.GCState{ + KeyspaceScope: &pdpb.KeyspaceScope{Keyspace: &pdpb.KeyspaceScope_KeyspaceId{KeyspaceId: 7}}, + TxnSafePoint: 10, + GcSafePoint: 5, + }}} + base := (&pdpb.WatchGCStatesResponse{Header: grpcutil.WrapHeader()}).Size() + delta := (&pdpb.WatchGCStatesResponse{Changes: []*pdpb.GCStateChange{change}}).Size() + + exact := splitWatchGCStatesResponses([]*pdpb.GCStateChange{change, change}, base+2*delta) + require.Len(t, exact, 1) + require.LessOrEqual(t, exact[0].Size(), base+2*delta) + + split := splitWatchGCStatesResponses([]*pdpb.GCStateChange{change, change}, base+2*delta-1) + require.Len(t, split, 2) + for _, response := range split { + require.NotNil(t, response.GetHeader()) + require.NotEmpty(t, response.GetChanges()) + require.LessOrEqual(t, response.Size(), base+2*delta-1) + } + + oversized := splitWatchGCStatesResponses([]*pdpb.GCStateChange{change}, base+delta-1) + require.Len(t, oversized, 1) + require.Greater(t, oversized[0].Size(), base+delta-1) + require.Empty(t, splitWatchGCStatesResponses(nil, base+delta)) +} + +type fakeGCStateChangeReceiver struct { + batches [][]gc.GCStateChange + receiveErr error + terminalErr error + receivedMaxes []int +} + +func (r *fakeGCStateChangeReceiver) RecvBatch(maxChanges int) ([]gc.GCStateChange, error) { + r.receivedMaxes = append(r.receivedMaxes, maxChanges) + if len(r.batches) == 0 { + return nil, r.receiveErr + } + batch := r.batches[0] + r.batches = r.batches[1:] + return batch, nil +} + +func (*fakeGCStateChangeReceiver) Done() <-chan struct{} { return nil } + +func (r *fakeGCStateChangeReceiver) Err() error { + return r.terminalErr +} + +type fakeWatchGCStatesServer struct { + ctx context.Context + sent []*pdpb.WatchGCStatesResponse + sendHook func(*pdpb.WatchGCStatesResponse) error +} + +func (s *fakeWatchGCStatesServer) Send(response *pdpb.WatchGCStatesResponse) error { + if s.sendHook != nil { + if err := s.sendHook(response); err != nil { + return err + } + } + s.sent = append(s.sent, response) + return nil +} + +func (*fakeWatchGCStatesServer) SetHeader(metadata.MD) error { return nil } +func (*fakeWatchGCStatesServer) SendHeader(metadata.MD) error { return nil } +func (*fakeWatchGCStatesServer) SetTrailer(metadata.MD) {} + +func (s *fakeWatchGCStatesServer) Context() context.Context { + if s.ctx == nil { + return context.Background() + } + return s.ctx +} + +func (*fakeWatchGCStatesServer) SendMsg(any) error { return nil } +func (*fakeWatchGCStatesServer) RecvMsg(any) error { return nil } + +func TestServeWatchGCStatesRechecksTerminalCauseBeforeEverySend(t *testing.T) { + state := gc.GCState{KeyspaceID: 7, TxnSafePoint: 10, GCSafePoint: 5} + receiver := &fakeGCStateChangeReceiver{ + batches: [][]gc.GCStateChange{{gc.NewGCStateUpsert(state), gc.NewGCStateUpsert(state)}}, + } + stream := &fakeWatchGCStatesServer{} + stream.sendHook = func(*pdpb.WatchGCStatesResponse) error { + receiver.terminalErr = errs.ErrNotLeader + return nil + } + protoChange := &pdpb.GCStateChange{Change: &pdpb.GCStateChange_Upsert{Upsert: &pdpb.GCState{ + KeyspaceScope: &pdpb.KeyspaceScope{Keyspace: &pdpb.KeyspaceScope_KeyspaceId{KeyspaceId: 7}}, + TxnSafePoint: 10, + GcSafePoint: 5, + }}} + maxSize := (&pdpb.WatchGCStatesResponse{Header: grpcutil.WrapHeader()}).Size() + + (&pdpb.WatchGCStatesResponse{Changes: []*pdpb.GCStateChange{protoChange}}).Size() + + err := serveWatchGCStates(receiver, stream, maxSize) + require.ErrorIs(t, err, errs.ErrNotLeader) + require.Equal(t, codes.Unavailable, status.Code(err)) + require.Len(t, stream.sent, 1) + require.Equal(t, []int{1024}, receiver.receivedMaxes) +} + +func TestServeWatchGCStatesRejectsInvalidInternalChange(t *testing.T) { + receiver := &fakeGCStateChangeReceiver{batches: [][]gc.GCStateChange{{{}}}} + stream := &fakeWatchGCStatesServer{} + + err := serveWatchGCStates(receiver, stream, 1024) + require.Equal(t, codes.Internal, status.Code(err)) + require.Empty(t, stream.sent) +} + +func TestServeWatchGCStatesReturnsRawSendError(t *testing.T) { + sendErr := errors.New("send failed") + receiver := &fakeGCStateChangeReceiver{ + batches: [][]gc.GCStateChange{{gc.NewGCStateRemoved(9)}}, + } + stream := &fakeWatchGCStatesServer{sendHook: func(*pdpb.WatchGCStatesResponse) error { + return sendErr + }} + + err := serveWatchGCStates(receiver, stream, 1024) + require.Same(t, sendErr, err) +} + +func TestWatchGCStatesErrorToStatus(t *testing.T) { + testCases := []struct { + name string + err error + code codes.Code + }{ + {name: "not leader", err: errs.ErrNotLeader, code: codes.Unavailable}, + {name: "initialization failure", err: errors.New("load initial GC states"), code: codes.Unavailable}, + {name: "slow consumer", err: errs.ErrGCStateWatcherSlowConsumer, code: codes.ResourceExhausted}, + {name: "canceled", err: context.Canceled, code: codes.Canceled}, + {name: "deadline exceeded", err: context.DeadlineExceeded, code: codes.DeadlineExceeded}, + } + + for _, testCase := range testCases { + t.Run(testCase.name, func(t *testing.T) { + err := watchGCStatesErrorToStatus(testCase.err) + require.Equal(t, testCase.code, status.Code(err)) + }) + } +} + +// cancelableGCStateReceiver keeps batch state in the sending worker and exposes +// the terminal cause through a context, which the supervisor can read safely. +type cancelableGCStateReceiver struct { + ctx context.Context + changes []gc.GCStateChange + receiveStarted chan struct{} + receiveExited chan struct{} +} + +func (r *cancelableGCStateReceiver) Done() <-chan struct{} { return r.ctx.Done() } +func (r *cancelableGCStateReceiver) Err() error { return context.Cause(r.ctx) } +func (r *cancelableGCStateReceiver) RecvBatch(maxChanges int) ([]gc.GCStateChange, error) { + if err := r.Err(); err != nil { + return nil, err + } + if len(r.changes) == 0 { + if r.receiveStarted != nil { + close(r.receiveStarted) + defer close(r.receiveExited) + } + <-r.Done() + return nil, r.Err() + } + n := min(maxChanges, len(r.changes)) + batch := r.changes[:n] + r.changes = r.changes[n:] + return batch, nil +} + +func waitWatchGCStatesSignal(t *testing.T, signal <-chan struct{}, message string) { + t.Helper() + select { + case <-signal: + case <-time.After(5 * time.Second): + require.FailNow(t, message) + } +} + +func TestServeWatchGCStatesCancellationUnblocksHandler(t *testing.T) { + for _, tc := range []struct { + name string + cause error + code codes.Code + }{ + {"leader loss", errs.ErrNotLeader, codes.Unavailable}, + {"slow consumer", errs.ErrGCStateWatcherSlowConsumer, codes.ResourceExhausted}, + } { + t.Run(tc.name, func(t *testing.T) { + streamCtx, cancelStream := context.WithCancel(context.Background()) + receiverCtx, cancelReceiver := context.WithCancelCause(streamCtx) + receiver := &cancelableGCStateReceiver{ctx: receiverCtx, changes: []gc.GCStateChange{gc.NewGCStateUpsert(gc.GCState{KeyspaceID: 7, TxnSafePoint: 10})}} + sendStarted, sendExited := make(chan struct{}), make(chan struct{}) + stream := &fakeWatchGCStatesServer{ctx: streamCtx, sendHook: func(*pdpb.WatchGCStatesResponse) error { + close(sendStarted) + defer close(sendExited) + <-streamCtx.Done() + return streamCtx.Err() + }} + handlerDone := make(chan struct{}) + var handlerErr error + t.Cleanup(func() { + cancelReceiver(context.Canceled) + cancelStream() + waitWatchGCStatesSignal(t, handlerDone, "handler did not clean up") + select { + case <-sendStarted: + waitWatchGCStatesSignal(t, sendExited, "send did not clean up") + default: + } + }) + go func() { defer close(handlerDone); handlerErr = serveWatchGCStates(receiver, stream, 1024) }() + waitWatchGCStatesSignal(t, sendStarted, "send did not start") + cancelReceiver(tc.cause) + waitWatchGCStatesSignal(t, handlerDone, "handler did not return while send was blocked") + require.Equal(t, tc.code, status.Code(handlerErr)) + require.NoError(t, streamCtx.Err()) + select { + case <-sendExited: + require.FailNow(t, "send exited before transport teardown") + default: + } + cancelStream() + waitWatchGCStatesSignal(t, sendExited, "send did not exit after transport teardown") + }) + } +} + +type observedWatchGCStatesStream struct { + pdpb.PD_WatchGCStatesServer + sentWireBytes int + observed bool + blockedSendStarted chan struct{} + blockedSendExited chan struct{} +} + +func (s *observedWatchGCStatesStream) Send(response *pdpb.WatchGCStatesResponse) error { + // grpc-go v1.82.1 starts with 64 KiB write quota. With the client's static + // 64 KiB receive window and no Recv calls, this next send cannot regain quota. + const receiveWindow = 64 << 10 + const writeQuota = 64 << 10 + if !s.observed && s.sentWireBytes >= receiveWindow+writeQuota { + s.observed = true + close(s.blockedSendStarted) + defer close(s.blockedSendExited) + } + err := s.PD_WatchGCStatesServer.Send(response) + if err == nil { + s.sentWireBytes += response.Size() + 5 + } + return err +} + +type watchGCStatesTransportServer struct { + pdpb.UnimplementedPDServer + changes []gc.GCStateChange + cancelReceiver chan context.CancelCauseFunc + handlerResult chan error + blockedSendStarted chan struct{} + blockedSendExited chan struct{} +} + +func (s *watchGCStatesTransportServer) WatchGCStates(_ *pdpb.WatchGCStatesRequest, stream pdpb.PD_WatchGCStatesServer) error { + ctx, cancel := context.WithCancelCause(stream.Context()) + defer cancel(context.Canceled) + s.cancelReceiver <- cancel + receiver := &cancelableGCStateReceiver{ctx: ctx, changes: s.changes} + observed := &observedWatchGCStatesStream{PD_WatchGCStatesServer: stream, blockedSendStarted: s.blockedSendStarted, blockedSendExited: s.blockedSendExited} + err := serveWatchGCStates(receiver, observed, maxWatchGCStatesResponseSize) + s.handlerResult <- err + return err +} + +func TestWatchGCStatesTransportCancellationUnblocksSend(t *testing.T) { + for _, tc := range []struct { + name string + cause error + code codes.Code + }{ + {"leader loss", errs.ErrNotLeader, codes.Unavailable}, + {"slow consumer", errs.ErrGCStateWatcherSlowConsumer, codes.ResourceExhausted}, + } { + t.Run(tc.name, func(t *testing.T) { + service := &watchGCStatesTransportServer{ + cancelReceiver: make(chan context.CancelCauseFunc, 1), handlerResult: make(chan error, 1), + blockedSendStarted: make(chan struct{}), blockedSendExited: make(chan struct{}), + } + for i := range 16 * 1024 { + service.changes = append(service.changes, gc.NewGCStateUpsert(gc.GCState{ + KeyspaceID: uint32(i), IsKeyspaceLevel: true, TxnSafePoint: math.MaxUint64, GCSafePoint: math.MaxUint64 - 1, + })) + } + listener := bufconn.Listen(1 << 20) + transport := grpc.NewServer() + pdpb.RegisterPDServer(transport, service) + serveErr := make(chan error, 1) + go func() { serveErr <- transport.Serve(listener) }() + t.Cleanup(func() { transport.Stop(); require.NoError(t, <-serveErr) }) + conn, err := grpc.NewClient("passthrough:///bufnet", + grpc.WithContextDialer(func(context.Context, string) (net.Conn, error) { return listener.Dial() }), + grpc.WithTransportCredentials(insecure.NewCredentials()), + grpc.WithStaticStreamWindowSize(64<<10), grpc.WithStaticConnWindowSize(64<<10), + ) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, conn.Close()) }) + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + t.Cleanup(cancel) + stream, err := pdpb.NewPDClient(conn).WatchGCStates(ctx, &pdpb.WatchGCStatesRequest{}) + require.NoError(t, err) + waitWatchGCStatesSignal(t, service.blockedSendStarted, "transport send did not reach exhausted quota") + select { + case <-service.blockedSendExited: + require.FailNow(t, "transport send unexpectedly completed") + default: + } + cancelReceiver := <-service.cancelReceiver + cancelReceiver(tc.cause) + select { + case err := <-service.handlerResult: + require.Equal(t, tc.code, status.Code(err)) + case <-time.After(5 * time.Second): + require.FailNow(t, "handler did not return while transport send was blocked") + } + waitWatchGCStatesSignal(t, service.blockedSendExited, "transport teardown did not unblock send") + require.NoError(t, ctx.Err()) + for err == nil { + _, err = stream.Recv() + } + require.Equal(t, tc.code, status.Code(err)) + }) + } +} + +func TestServeWatchGCStatesCancellationBeforeServing(t *testing.T) { + ctx, cancel := context.WithCancelCause(context.Background()) + cancel(errs.ErrNotLeader) + receiver := &cancelableGCStateReceiver{ctx: ctx, changes: []gc.GCStateChange{gc.NewGCStateRemoved(7)}} + stream := &fakeWatchGCStatesServer{} + err := serveWatchGCStates(receiver, stream, 1024) + require.Equal(t, codes.Unavailable, status.Code(err)) + require.Empty(t, stream.sent) +} + +func TestServeWatchGCStatesCancellationWhileReceiving(t *testing.T) { + for _, parentCanceled := range []bool{false, true} { + name := "watcher" + if parentCanceled { + name = "parent stream" + } + t.Run(name, func(t *testing.T) { + streamCtx, cancelStream := context.WithCancel(context.Background()) + ctx, cancelReceiver := context.WithCancelCause(streamCtx) + receiver := &cancelableGCStateReceiver{ctx: ctx, receiveStarted: make(chan struct{}), receiveExited: make(chan struct{})} + stream := &fakeWatchGCStatesServer{ctx: streamCtx} + handlerDone := make(chan struct{}) + var handlerErr error + t.Cleanup(func() { + cancelReceiver(context.Canceled) + cancelStream() + waitWatchGCStatesSignal(t, handlerDone, "handler did not clean up") + select { + case <-receiver.receiveStarted: + waitWatchGCStatesSignal(t, receiver.receiveExited, "receive did not clean up") + default: + } + }) + go func() { defer close(handlerDone); handlerErr = serveWatchGCStates(receiver, stream, 1024) }() + waitWatchGCStatesSignal(t, receiver.receiveStarted, "receive did not start") + want := codes.Unavailable + if parentCanceled { + cancelStream() + want = codes.Canceled + } else { + cancelReceiver(errs.ErrNotLeader) + } + waitWatchGCStatesSignal(t, handlerDone, "handler did not return after cancellation") + waitWatchGCStatesSignal(t, receiver.receiveExited, "receive did not return after cancellation") + require.Equal(t, want, status.Code(handlerErr)) + require.Empty(t, stream.sent) + }) + } +} diff --git a/server/server.go b/server/server.go index 20e0579dea9..a3ba0ce3bf8 100644 --- a/server/server.go +++ b/server/server.go @@ -90,6 +90,7 @@ import ( "github.com/tikv/pd/pkg/utils/tsoutil" "github.com/tikv/pd/pkg/utils/typeutil" "github.com/tikv/pd/pkg/versioninfo" + "github.com/tikv/pd/pkg/versioninfo/kerneltype" "github.com/tikv/pd/server/cluster" "github.com/tikv/pd/server/config" @@ -563,6 +564,9 @@ func (s *Server) startServer(ctx context.Context) error { log.Info("no metering config provided, the metering writer will not be started") } s.gcStateManager = gc.NewGCStateManager(s.storage.GetGCStateProvider(), s.cfg.PDServerCfg, s.keyspaceManager) + if kerneltype.IsNextGen() { + s.gcStateManager.SetEtcdClient(s.client) + } s.hbStreams = hbstream.NewHeartbeatStreams(ctx, "", s.cluster) // initial hot_region_storage in here. diff --git a/tests/integrations/go.mod b/tests/integrations/go.mod index 00b9e075b4f..eb933b245ba 100644 --- a/tests/integrations/go.mod +++ b/tests/integrations/go.mod @@ -15,7 +15,7 @@ require ( github.com/golang/protobuf v1.5.4 github.com/pingcap/errors v0.11.5-0.20211224045212-9687c2b0f87c github.com/pingcap/failpoint v0.0.0-20240528011301-b51a646c7c86 - github.com/pingcap/kvproto v0.0.0-20260903054228-107095f1d250 + github.com/pingcap/kvproto v0.0.0-20260903062353-65b4e27a438d github.com/pingcap/log v1.1.1-0.20221110025148-ca232912c9f3 github.com/prometheus/client_golang v1.20.5 github.com/prometheus/client_model v0.6.1 diff --git a/tests/integrations/go.sum b/tests/integrations/go.sum index 5db28b57234..2a17320ad60 100644 --- a/tests/integrations/go.sum +++ b/tests/integrations/go.sum @@ -483,8 +483,8 @@ github.com/pingcap/errors v0.11.5-0.20211224045212-9687c2b0f87c/go.mod h1:X2r9ue github.com/pingcap/failpoint v0.0.0-20240528011301-b51a646c7c86 h1:tdMsjOqUR7YXHoBitzdebTvOjs/swniBTOLy5XiMtuE= github.com/pingcap/failpoint v0.0.0-20240528011301-b51a646c7c86/go.mod h1:exzhVYca3WRtd6gclGNErRWb1qEgff3LYta0LvRmON4= github.com/pingcap/kvproto v0.0.0-20191211054548-3c6b38ea5107/go.mod h1:WWLmULLO7l8IOcQG+t+ItJ3fEcrL5FxF0Wu+HrMy26w= -github.com/pingcap/kvproto v0.0.0-20260903054228-107095f1d250 h1:6yUryXKVbKpCNdZWL58/OcZj8NPLUA/xsJYXSbsD59w= -github.com/pingcap/kvproto v0.0.0-20260903054228-107095f1d250/go.mod h1:z6+aAHB7dBkA+LyinEX+48/ImRJ3jag0Hg0c7wkhEvE= +github.com/pingcap/kvproto v0.0.0-20260903062353-65b4e27a438d h1:KS3ekak/ljCj5xvkGqbwVLi2eL7B8GFSYzU9TOUUPPo= +github.com/pingcap/kvproto v0.0.0-20260903062353-65b4e27a438d/go.mod h1:z6+aAHB7dBkA+LyinEX+48/ImRJ3jag0Hg0c7wkhEvE= github.com/pingcap/log v0.0.0-20210625125904-98ed8e2eb1c7/go.mod h1:8AanEdAHATuRurdGxZXBz0At+9avep+ub7U1AGYLIMM= github.com/pingcap/log v1.1.1-0.20221110025148-ca232912c9f3 h1:HR/ylkkLmGdSSDaD8IDP+SZrdhV1Kibl9KrHxJ9eciw= github.com/pingcap/log v1.1.1-0.20221110025148-ca232912c9f3/go.mod h1:DWQW5jICDR7UJh4HtxXSM20Churx4CQL0fwL/SoOSA4= diff --git a/tests/server/gc/gc_test.go b/tests/server/gc/gc_test.go index d1d261ae203..66601e94afd 100644 --- a/tests/server/gc/gc_test.go +++ b/tests/server/gc/gc_test.go @@ -16,6 +16,7 @@ package gc import ( "context" + "errors" "math" "slices" "sync" @@ -23,16 +24,22 @@ import ( "testing" "time" + "github.com/prometheus/client_golang/prometheus" "github.com/stretchr/testify/require" "go.uber.org/goleak" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/metadata" + "google.golang.org/grpc/status" "github.com/pingcap/failpoint" "github.com/pingcap/kvproto/pkg/pdpb" "github.com/tikv/pd/pkg/keyspace" "github.com/tikv/pd/pkg/keyspace/constant" + "github.com/tikv/pd/pkg/ratelimit" "github.com/tikv/pd/pkg/utils/testutil" "github.com/tikv/pd/pkg/versioninfo/kerneltype" + "github.com/tikv/pd/server" "github.com/tikv/pd/server/config" "github.com/tikv/pd/tests" ) @@ -48,6 +55,7 @@ const ( postGetGCStateCallFailpoint = "github.com/tikv/pd/server/postGetGCStateCall" getGCStateBeforeSlowPathFailpoint = "github.com/tikv/pd/pkg/gc/getGCStateBeforeSlowPath" skipCampaignLeaderCheckFailpoint = "github.com/tikv/pd/pkg/member/skipCampaignLeaderCheck" + watchGCStatesRegisteredFailpoint = "github.com/tikv/pd/pkg/gc/watchGCStatesRegistered" ) func makeKeyspaceScope(keyspaceID uint32) *pdpb.KeyspaceScope { @@ -139,6 +147,179 @@ func (p *blockingFailpoint) releaseAndDisable(re *require.Assertions) { }) } +type watchGCStatesRegistrationPoint struct { + registered chan struct{} + registerOnce sync.Once + disableOnce sync.Once +} + +func enableWatchGCStatesRegistrationPoint(t *testing.T) *watchGCStatesRegistrationPoint { + t.Helper() + re := require.New(t) + point := &watchGCStatesRegistrationPoint{registered: make(chan struct{})} + re.NoError(failpoint.EnableCall(watchGCStatesRegisteredFailpoint, func() { + point.registerOnce.Do(func() { + close(point.registered) + }) + })) + t.Cleanup(func() { + point.disable(re) + }) + return point +} + +func (p *watchGCStatesRegistrationPoint) wait(t *testing.T) { + t.Helper() + select { + case <-p.registered: + case <-time.After(5 * time.Second): + require.FailNow(t, "WatchGCStates was not registered") + } +} + +func (p *watchGCStatesRegistrationPoint) disable(re *require.Assertions) { + p.disableOnce.Do(func() { + re.NoError(failpoint.Disable(watchGCStatesRegisteredFailpoint)) + }) +} + +func newWatchGCStatesCluster(t *testing.T, serverCount int, bootstrap bool) *tests.TestCluster { + t.Helper() + re := require.New(t) + ctx, cancel := context.WithCancel(context.Background()) + cluster, err := tests.NewTestCluster(ctx, serverCount, func(conf *config.Config, _ string) { + conf.Keyspace.WaitRegionSplit = false + }) + re.NoError(err) + t.Cleanup(func() { + cancel() + cluster.Destroy() + }) + re.NoError(cluster.RunInitialServers()) + re.NotEmpty(cluster.WaitLeader()) + if bootstrap { + re.NoError(cluster.GetLeaderServer().BootstrapCluster()) + } + return cluster +} + +func newWatchGCStatesClient(t *testing.T, addr string) pdpb.PDClient { + t.Helper() + re := require.New(t) + client, conn := testutil.MustNewGrpcClient(re, addr) + t.Cleanup(func() { + re.NoError(conn.Close()) + }) + return client +} + +func openWatchGCStates( + t *testing.T, + client pdpb.PDClient, + header *pdpb.RequestHeader, + skipLoadingInitial bool, +) (pdpb.PD_WatchGCStatesClient, context.CancelFunc) { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 20*time.Second) + t.Cleanup(cancel) + stream, err := client.WatchGCStates(ctx, &pdpb.WatchGCStatesRequest{ + Header: header, + SkipLoadingInitial: skipLoadingInitial, + }) + require.NoError(t, err) + return stream, cancel +} + +func recvWatchGCStateForKeyspace(t *testing.T, stream pdpb.PD_WatchGCStatesClient, keyspaceID uint32) *pdpb.GCState { + t.Helper() + for { + response, err := stream.Recv() + require.NoError(t, err) + require.NotNil(t, response.GetHeader()) + for _, change := range response.GetChanges() { + if state := change.GetUpsert(); state != nil && state.GetKeyspaceScope().GetKeyspaceId() == keyspaceID { + return state + } + } + } +} + +func advanceWatchGCStatesTxnSafePoint( + t *testing.T, + client pdpb.PDClient, + header *pdpb.RequestHeader, + keyspaceID uint32, + target uint64, +) { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 20*time.Second) + defer cancel() + response, err := client.AdvanceTxnSafePoint(ctx, &pdpb.AdvanceTxnSafePointRequest{ + Header: header, + KeyspaceScope: makeKeyspaceScope(keyspaceID), + Target: target, + }) + require.NoError(t, err) + require.NotNil(t, response.GetHeader()) + require.Nil(t, response.GetHeader().GetError()) + require.Equal(t, target, response.GetNewTxnSafePoint()) +} + +type failingWatchGCStatesServer struct { + ctx context.Context + sendErr error +} + +func (s *failingWatchGCStatesServer) Send(*pdpb.WatchGCStatesResponse) error { + return s.sendErr +} + +func (*failingWatchGCStatesServer) SetHeader(metadata.MD) error { return nil } +func (*failingWatchGCStatesServer) SendHeader(metadata.MD) error { return nil } +func (*failingWatchGCStatesServer) SetTrailer(metadata.MD) {} + +func (s *failingWatchGCStatesServer) Context() context.Context { + return s.ctx +} + +func (*failingWatchGCStatesServer) SendMsg(any) error { return nil } +func (*failingWatchGCStatesServer) RecvMsg(any) error { return nil } + +func prometheusMetricValue(t *testing.T, name string, labels map[string]string) float64 { + t.Helper() + families, err := prometheus.DefaultGatherer.Gather() + require.NoError(t, err) + for _, family := range families { + if family.GetName() != name { + continue + } + for _, metric := range family.GetMetric() { + if len(metric.GetLabel()) != len(labels) { + continue + } + matches := true + for _, pair := range metric.GetLabel() { + if labels[pair.GetName()] != pair.GetValue() { + matches = false + break + } + } + if !matches { + continue + } + if gauge := metric.GetGauge(); gauge != nil { + return gauge.GetValue() + } + if counter := metric.GetCounter(); counter != nil { + return counter.GetValue() + } + require.FailNow(t, "metric has unsupported type", name) + } + } + require.FailNow(t, "metric not found", name) + return 0 +} + func TestGCOperations(t *testing.T) { re := require.New(t) ctx, cancel := context.WithCancel(context.Background()) @@ -883,3 +1064,386 @@ func TestGetGCStateSlowPathReadsLatestStateIfLeaderLostBeforeRead(t *testing.T) re.Nil(res.resp.GetHeader().GetError()) re.Equal(uint64(20), res.resp.GetGcState().GetTxnSafePoint()) } + +func TestWatchGCStatesInitialAndSkipInitialRegistrationBoundary(t *testing.T) { + re := require.New(t) + cluster := newWatchGCStatesCluster(t, 1, true) + leaderServer := cluster.GetLeaderServer() + re.NotNil(leaderServer) + + ks, err := leaderServer.GetKeyspaceManager().CreateKeyspace(&keyspace.CreateKeyspaceRequest{ + Name: "watch-gc-states", + Config: map[string]string{keyspace.GCManagementType: keyspace.KeyspaceLevelGC}, + CreateTime: time.Now().Unix(), + }) + re.NoError(err) + + client := newWatchGCStatesClient(t, leaderServer.GetAddr()) + header := testutil.NewRequestHeader(leaderServer.GetClusterID()) + initialStream, cancelInitial := openWatchGCStates(t, client, header, false) + initial := recvWatchGCStateForKeyspace(t, initialStream, ks.GetId()) + re.True(initial.GetIsKeyspaceLevelGc()) + re.Zero(initial.GetTxnSafePoint()) + re.Zero(initial.GetGcSafePoint()) + re.Empty(initial.GetGcBarriers()) + + advanceWatchGCStatesTxnSafePoint(t, client, header, ks.GetId(), 10) + live := recvWatchGCStateForKeyspace(t, initialStream, ks.GetId()) + re.True(live.GetIsKeyspaceLevelGc()) + re.Equal(uint64(10), live.GetTxnSafePoint()) + re.Zero(live.GetGcSafePoint()) + re.Empty(live.GetGcBarriers()) + cancelInitial() + + registration := enableWatchGCStatesRegistrationPoint(t) + skipInitialStream, _ := openWatchGCStates(t, client, header, true) + registration.wait(t) + registration.disable(re) + + advanceWatchGCStatesTxnSafePoint(t, client, header, ks.GetId(), 20) + firstResponse, err := skipInitialStream.Recv() + re.NoError(err) + re.NotNil(firstResponse.GetHeader()) + re.Len(firstResponse.GetChanges(), 1) + firstAfterRegistration := firstResponse.GetChanges()[0].GetUpsert() + re.NotNil(firstAfterRegistration) + re.Equal(ks.GetId(), firstAfterRegistration.GetKeyspaceScope().GetKeyspaceId()) + re.True(firstAfterRegistration.GetIsKeyspaceLevelGc()) + re.Equal(uint64(20), firstAfterRegistration.GetTxnSafePoint()) + re.Zero(firstAfterRegistration.GetGcSafePoint()) + re.Empty(firstAfterRegistration.GetGcBarriers()) +} + +func TestWatchGCStatesUsesIndexOnlyInNextGen(t *testing.T) { + cluster := newWatchGCStatesCluster(t, 1, true) + leader := cluster.GetLeaderServer() + require.NotNil(t, leader) + const legacyIterator = "github.com/tikv/pd/pkg/gc/iterateAllKeyspacesGCStatesError" + const legacyError = "legacy keyspace iterator reached" + require.NoError(t, failpoint.Enable(legacyIterator, `return("legacy keyspace iterator reached")`)) + defer func() { require.NoError(t, failpoint.Disable(legacyIterator)) }() + + client := newWatchGCStatesClient(t, leader.GetAddr()) + stream, _ := openWatchGCStates(t, client, testutil.NewRequestHeader(leader.GetClusterID()), false) + if kerneltype.IsNextGen() { + state := recvWatchGCStateForKeyspace(t, stream, constant.NullKeyspaceID) + require.False(t, state.GetIsKeyspaceLevelGc()) + return + } + response, err := stream.Recv() + require.Nil(t, response) + require.ErrorContains(t, err, legacyError) + require.Equal(t, codes.Unavailable, status.Code(err)) +} + +func TestWatchGCStatesRequestPreflight(t *testing.T) { + tests := []struct { + name string + setup func(*testing.T) (string, *pdpb.RequestHeader) + wantCode codes.Code + }{ + { + name: "wrong cluster ID", + setup: func(t *testing.T) (string, *pdpb.RequestHeader) { + cluster := newWatchGCStatesCluster(t, 1, true) + leader := cluster.GetLeaderServer() + return leader.GetAddr(), testutil.NewRequestHeader(leader.GetClusterID() + 1) + }, + wantCode: codes.FailedPrecondition, + }, + { + name: "direct follower", + setup: func(t *testing.T) (string, *pdpb.RequestHeader) { + cluster := newWatchGCStatesCluster(t, 2, true) + follower := cluster.GetServer(cluster.GetFollower()) + require.NotNil(t, follower) + return follower.GetAddr(), testutil.NewRequestHeader(follower.GetClusterID()) + }, + wantCode: codes.Unavailable, + }, + { + name: "unbootstrapped leader", + setup: func(t *testing.T) (string, *pdpb.RequestHeader) { + cluster := newWatchGCStatesCluster(t, 1, false) + leader := cluster.GetLeaderServer() + return leader.GetAddr(), testutil.NewRequestHeader(leader.GetClusterID()) + }, + wantCode: codes.Unavailable, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + addr, header := test.setup(t) + client := newWatchGCStatesClient(t, addr) + stream, _ := openWatchGCStates(t, client, header, true) + response, err := stream.Recv() + require.Nil(t, response) + require.Equal(t, test.wantCode, status.Code(err)) + }) + } +} + +func TestWatchGCStatesSendFailureCleansUpPublicHandler(t *testing.T) { + re := require.New(t) + cluster := newWatchGCStatesCluster(t, 1, true) + leaderServer := cluster.GetLeaderServer() + re.NotNil(leaderServer) + pdServer := leaderServer.GetServer() + + limiter := limitWatchGCStatesConcurrency(t, pdServer) + + activeBefore := prometheusMetricValue(t, "pd_gc_watcher_count", nil) + clientCancelBefore := prometheusMetricValue(t, "pd_gc_watcher_termination_total", map[string]string{"reason": "client_cancel"}) + registration := enableWatchGCStatesRegistrationPoint(t) + sendErr := errors.New("send failed") + streamCtx, cancelStream := context.WithTimeout(context.Background(), 20*time.Second) + defer cancelStream() + stream := &failingWatchGCStatesServer{ctx: streamCtx, sendErr: sendErr} + handlerDone := make(chan error, 1) + go func() { + handlerDone <- (&server.GrpcServer{Server: pdServer}).WatchGCStates(&pdpb.WatchGCStatesRequest{ + Header: testutil.NewRequestHeader(leaderServer.GetClusterID()), + SkipLoadingInitial: true, + }, stream) + }() + + registration.wait(t) + registration.disable(re) + re.Equal(activeBefore+1, prometheusMetricValue(t, "pd_gc_watcher_count", nil)) + limit, current := limiter.GetConcurrencyLimiterStatus("WatchGCStates") + re.Equal(uint64(1), limit) + re.Equal(uint64(1), current) + + _, err := pdServer.GetGCStateManager().AdvanceTxnSafePoint(constant.NullKeyspaceID, 10, time.Now()) + re.NoError(err) + select { + case err := <-handlerDone: + re.Same(sendErr, err) + case <-time.After(5 * time.Second): + re.FailNow("WatchGCStates handler did not return after the send failure") + } + + re.Equal(activeBefore, prometheusMetricValue(t, "pd_gc_watcher_count", nil)) + re.Equal(clientCancelBefore+1, prometheusMetricValue(t, "pd_gc_watcher_termination_total", map[string]string{"reason": "client_cancel"})) + _, current = limiter.GetConcurrencyLimiterStatus("WatchGCStates") + re.Zero(current) +} + +func TestWatchGCStatesHoldsRateLimitTokenForStreamLifetime(t *testing.T) { + re := require.New(t) + cluster := newWatchGCStatesCluster(t, 1, true) + leaderServer := cluster.GetLeaderServer() + re.NotNil(leaderServer) + limiter := limitWatchGCStatesConcurrency(t, leaderServer.GetServer()) + + client := newWatchGCStatesClient(t, leaderServer.GetAddr()) + header := testutil.NewRequestHeader(leaderServer.GetClusterID()) + firstRegistration := enableWatchGCStatesRegistrationPoint(t) + _, cancelFirst := openWatchGCStates(t, client, header, true) + firstRegistration.wait(t) + firstRegistration.disable(re) + limit, current := limiter.GetConcurrencyLimiterStatus("WatchGCStates") + re.Equal(uint64(1), limit) + re.Equal(uint64(1), current) + + secondStream, _ := openWatchGCStates(t, client, header, true) + response, err := secondStream.Recv() + re.Nil(response) + re.Equal(codes.ResourceExhausted, status.Code(err)) + + cancelFirst() + testutil.Eventually(re, func() bool { + _, current := limiter.GetConcurrencyLimiterStatus("WatchGCStates") + return current == 0 + }, testutil.WithWaitFor(5*time.Second), testutil.WithTickInterval(10*time.Millisecond)) + + thirdRegistration := enableWatchGCStatesRegistrationPoint(t) + thirdStream, _ := openWatchGCStates(t, client, header, true) + thirdRegistration.wait(t) + thirdRegistration.disable(re) + advanceWatchGCStatesTxnSafePoint(t, client, header, constant.NullKeyspaceID, 10) + state := recvWatchGCStateForKeyspace(t, thirdStream, constant.NullKeyspaceID) + re.Equal(uint64(10), state.GetTxnSafePoint()) +} + +func TestWatchGCStatesTerminatesOnLeaderTransferAndReinitializes(t *testing.T) { + re := require.New(t) + cluster, req, cleanup := newGCStateLeaderTransitionCluster(t) + t.Cleanup(cleanup) + + oldLeader := cluster.GetLeader() + re.NotEmpty(oldLeader) + oldLeaderServer := cluster.GetServer(oldLeader) + re.NotNil(oldLeaderServer) + oldClient := newWatchGCStatesClient(t, oldLeaderServer.GetAddr()) + oldStream, _ := openWatchGCStates(t, oldClient, req.GetHeader(), false) + initial := recvWatchGCStateForKeyspace(t, oldStream, constant.NullKeyspaceID) + re.False(initial.GetIsKeyspaceLevelGc()) + re.Zero(initial.GetTxnSafePoint()) + re.Zero(initial.GetGcSafePoint()) + re.Empty(initial.GetGcBarriers()) + + re.NoError(oldLeaderServer.ResignLeaderWithRetry()) + newLeader := cluster.WaitLeader() + re.NotEmpty(newLeader) + re.NotEqual(oldLeader, newLeader) + for { + response, err := oldStream.Recv() + if err != nil { + re.Nil(response) + re.Equal(codes.Unavailable, status.Code(err)) + break + } + re.NotNil(response) + } + + newLeaderServer := cluster.GetServer(newLeader) + re.NotNil(newLeaderServer) + newClient := newWatchGCStatesClient(t, newLeaderServer.GetAddr()) + advanceWatchGCStatesTxnSafePoint(t, newClient, req.GetHeader(), constant.NullKeyspaceID, 10) + newStream, _ := openWatchGCStates(t, newClient, req.GetHeader(), false) + reinitialized := recvWatchGCStateForKeyspace(t, newStream, constant.NullKeyspaceID) + re.False(reinitialized.GetIsKeyspaceLevelGc()) + re.Equal(uint64(10), reinitialized.GetTxnSafePoint()) + re.Zero(reinitialized.GetGcSafePoint()) + re.Empty(reinitialized.GetGcBarriers()) +} + +func limitWatchGCStatesConcurrency(t *testing.T, pdServer *server.Server) *ratelimit.Controller { + t.Helper() + options := pdServer.GetServiceMiddlewarePersistOptions() + previousConfig := options.GetGRPCRateLimitConfig().Clone() + enabledConfig := previousConfig.Clone() + enabledConfig.EnableRateLimit = true + options.SetGRPCRateLimitConfig(enabledConfig) + limiter := pdServer.GetGRPCRateLimiter() + limiter.Update("WatchGCStates", ratelimit.UpdateConcurrencyLimiter(1)) + t.Cleanup(func() { + limiter.Update("WatchGCStates", ratelimit.UpdateConcurrencyLimiter(0)) + options.SetGRPCRateLimitConfig(previousConfig) + }) + return limiter +} + +type blockedWatchGCStatesServer struct { + failingWatchGCStatesServer + sendStarted chan struct{} + sendExited chan struct{} +} + +func (s *blockedWatchGCStatesServer) Send(*pdpb.WatchGCStatesResponse) error { + close(s.sendStarted) + defer close(s.sendExited) + <-s.ctx.Done() + return s.ctx.Err() +} + +func waitWatchGCStatesSignal(t *testing.T, signal <-chan struct{}, message string) { + t.Helper() + select { + case <-signal: + case <-time.After(5 * time.Second): + require.FailNow(t, message) + } +} + +func TestWatchGCStatesBlockedSendCleansUpPublicHandler(t *testing.T) { + for _, tc := range []struct { + name string + reason string + code codes.Code + }{ + {"leader loss", "leader_lost", codes.Unavailable}, + {"slow consumer", "slow_consumer", codes.ResourceExhausted}, + } { + t.Run(tc.name, func(t *testing.T) { + re := require.New(t) + cluster := newWatchGCStatesCluster(t, 1, true) + leaderServer := cluster.GetLeaderServer() + re.NotNil(leaderServer) + pdServer := leaderServer.GetServer() + manager := pdServer.GetGCStateManager() + limiter := limitWatchGCStatesConcurrency(t, pdServer) + activeBefore := prometheusMetricValue(t, "pd_gc_watcher_count", nil) + labels := map[string]string{"reason": tc.reason} + terminatedBefore := prometheusMetricValue(t, "pd_gc_watcher_termination_total", labels) + registration := enableWatchGCStatesRegistrationPoint(t) + streamCtx, cancelStream := context.WithTimeout(context.Background(), 30*time.Second) + stream := &blockedWatchGCStatesServer{ + failingWatchGCStatesServer: failingWatchGCStatesServer{ctx: streamCtx}, + sendStarted: make(chan struct{}), sendExited: make(chan struct{}), + } + handlerDone := make(chan struct{}) + var handlerErr error + t.Cleanup(func() { + cancelStream() + waitWatchGCStatesSignal(t, handlerDone, "public handler did not clean up") + select { + case <-stream.sendStarted: + waitWatchGCStatesSignal(t, stream.sendExited, "send did not clean up") + default: + } + }) + go func() { + defer close(handlerDone) + handlerErr = (&server.GrpcServer{Server: pdServer}).WatchGCStates(&pdpb.WatchGCStatesRequest{ + Header: testutil.NewRequestHeader(leaderServer.GetClusterID()), SkipLoadingInitial: true, + }, stream) + }() + registration.wait(t) + registration.disable(re) + re.Equal(activeBefore+1, prometheusMetricValue(t, "pd_gc_watcher_count", nil)) + _, current := limiter.GetConcurrencyLimiterStatus("WatchGCStates") + re.Equal(uint64(1), current) + res, err := manager.AdvanceTxnSafePoint(constant.NullKeyspaceID, 10, time.Now()) + re.NoError(err) + re.Equal(uint64(10), res.NewTxnSafePoint) + waitWatchGCStatesSignal(t, stream.sendStarted, "send did not start") + if tc.reason == "leader_lost" { + stop := manager.OnNodeBecomesLeader() + t.Cleanup(stop) + } else { + // The sending worker cannot receive these 1025 updates, overflowing + // the default live queue's 1024 slots. + for target := uint64(11); target <= 1035; target++ { + res, err := manager.AdvanceTxnSafePoint(constant.NullKeyspaceID, target, time.Now()) + re.NoError(err) + re.Equal(target, res.NewTxnSafePoint) + } + } + waitWatchGCStatesSignal(t, handlerDone, "public handler did not return while send was blocked") + re.Equal(tc.code, status.Code(handlerErr)) + re.NoError(streamCtx.Err()) + select { + case <-stream.sendExited: + re.FailNow("send exited before transport teardown") + default: + } + re.Equal(activeBefore, prometheusMetricValue(t, "pd_gc_watcher_count", nil)) + re.Equal(terminatedBefore+1, prometheusMetricValue(t, "pd_gc_watcher_termination_total", labels)) + _, current = limiter.GetConcurrencyLimiterStatus("WatchGCStates") + re.Zero(current) + + // Admission while the old send is still blocked proves the public + // handler released its token independently of transport progress. + nextRegistration := enableWatchGCStatesRegistrationPoint(t) + client := newWatchGCStatesClient(t, leaderServer.GetAddr()) + nextStream, cancelNext := openWatchGCStates(t, client, testutil.NewRequestHeader(leaderServer.GetClusterID()), true) + nextRegistration.wait(t) + nextRegistration.disable(re) + _, current = limiter.GetConcurrencyLimiterStatus("WatchGCStates") + re.Equal(uint64(1), current) + cancelStream() + waitWatchGCStatesSignal(t, stream.sendExited, "send did not exit after transport teardown") + cancelNext() + _, err = nextStream.Recv() + re.Equal(codes.Canceled, status.Code(err)) + testutil.Eventually(re, func() bool { + _, current := limiter.GetConcurrencyLimiterStatus("WatchGCStates") + return current == 0 && prometheusMetricValue(t, "pd_gc_watcher_count", nil) == activeBefore + }) + re.Equal(terminatedBefore+1, prometheusMetricValue(t, "pd_gc_watcher_termination_total", labels)) + }) + } +} diff --git a/tools/go.mod b/tools/go.mod index 747f9198561..7e0027895b5 100644 --- a/tools/go.mod +++ b/tools/go.mod @@ -23,7 +23,7 @@ require ( github.com/mattn/go-shellwords v1.0.12 github.com/pingcap/errors v0.11.5-0.20211224045212-9687c2b0f87c github.com/pingcap/failpoint v0.0.0-20240528011301-b51a646c7c86 - github.com/pingcap/kvproto v0.0.0-20260903054228-107095f1d250 + github.com/pingcap/kvproto v0.0.0-20260903062353-65b4e27a438d github.com/pingcap/log v1.1.1-0.20221110025148-ca232912c9f3 github.com/pmezard/go-difflib v1.0.0 github.com/prometheus/client_golang v1.20.5 diff --git a/tools/go.sum b/tools/go.sum index 769f9c9e60d..efdcb12f0d3 100644 --- a/tools/go.sum +++ b/tools/go.sum @@ -488,8 +488,8 @@ github.com/pingcap/errors v0.11.5-0.20211224045212-9687c2b0f87c/go.mod h1:X2r9ue github.com/pingcap/failpoint v0.0.0-20240528011301-b51a646c7c86 h1:tdMsjOqUR7YXHoBitzdebTvOjs/swniBTOLy5XiMtuE= github.com/pingcap/failpoint v0.0.0-20240528011301-b51a646c7c86/go.mod h1:exzhVYca3WRtd6gclGNErRWb1qEgff3LYta0LvRmON4= github.com/pingcap/kvproto v0.0.0-20191211054548-3c6b38ea5107/go.mod h1:WWLmULLO7l8IOcQG+t+ItJ3fEcrL5FxF0Wu+HrMy26w= -github.com/pingcap/kvproto v0.0.0-20260903054228-107095f1d250 h1:6yUryXKVbKpCNdZWL58/OcZj8NPLUA/xsJYXSbsD59w= -github.com/pingcap/kvproto v0.0.0-20260903054228-107095f1d250/go.mod h1:z6+aAHB7dBkA+LyinEX+48/ImRJ3jag0Hg0c7wkhEvE= +github.com/pingcap/kvproto v0.0.0-20260903062353-65b4e27a438d h1:KS3ekak/ljCj5xvkGqbwVLi2eL7B8GFSYzU9TOUUPPo= +github.com/pingcap/kvproto v0.0.0-20260903062353-65b4e27a438d/go.mod h1:z6+aAHB7dBkA+LyinEX+48/ImRJ3jag0Hg0c7wkhEvE= github.com/pingcap/log v0.0.0-20210625125904-98ed8e2eb1c7/go.mod h1:8AanEdAHATuRurdGxZXBz0At+9avep+ub7U1AGYLIMM= github.com/pingcap/log v1.1.1-0.20221110025148-ca232912c9f3 h1:HR/ylkkLmGdSSDaD8IDP+SZrdhV1Kibl9KrHxJ9eciw= github.com/pingcap/log v1.1.1-0.20221110025148-ca232912c9f3/go.mod h1:DWQW5jICDR7UJh4HtxXSM20Churx4CQL0fwL/SoOSA4=