diff --git a/logservice/logpuller/memory_quota.go b/logservice/logpuller/memory_quota.go new file mode 100644 index 0000000000..fbc0bc6c86 --- /dev/null +++ b/logservice/logpuller/memory_quota.go @@ -0,0 +1,386 @@ +// Copyright 2026 PingCAP, Inc. +// +// 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 logpuller + +import ( + "context" + "math" + "sync" + "sync/atomic" + "time" + + "github.com/pingcap/ticdc/pkg/metrics" + "github.com/tikv/client-go/v2/oracle" +) + +const ( + // defaultPauseLowPriorityRatio pauses new low-priority scans when memory + // pressure reaches 15% of the soft capacity. + defaultPauseLowPriorityRatio = 0.15 + + // defaultResumeLowPriorityRatio resumes low-priority scans after memory + // pressure falls to 5% of the soft capacity. + defaultResumeLowPriorityRatio = 0.05 + + // defaultHardLimitRatio blocks receiving more events when accounted event + // memory reaches twice the soft capacity. + defaultHardLimitRatio = 2.0 + + // defaultScanLagUnit is the lag unit used by the logarithmic scan estimate. + defaultScanLagUnit = 10 * time.Minute + + // defaultScanLagWeight controls how quickly the scan estimate grows with lag. + defaultScanLagWeight = 0.22 + + // defaultMaxScanLagFactor caps one scan estimate at this multiple of the base. + defaultMaxScanLagFactor = 16 +) + +type admissionLevel uint8 + +const ( + admissionNormal admissionLevel = iota + admissionPauseLowPriority +) + +// eventMemoryNotifier wakes event receivers that are waiting for memory. Each +// notification closes the current ready channel to wake all current waiters, +// then creates a new channel for future waiters. +// +// To avoid missing a notification, a receiver waits in this order: +// +// 1. Register the waiter. +// 2. Read the current ready channel under mu. +// 3. Recheck memory and the span state before blocking on that channel. +// +// If a notification happens just before registration, the final recheck sees +// the released memory or stopped span, so the receiver does not block. +type eventMemoryNotifier struct { + mu sync.Mutex + ready chan struct{} + waiters atomic.Int64 +} + +func newEventMemoryNotifier() *eventMemoryNotifier { + return &eventMemoryNotifier{ready: make(chan struct{})} +} + +func (n *eventMemoryNotifier) wait( + ctx context.Context, + span *subscribedSpan, + tryAcquire func() bool, +) bool { + n.waiters.Add(1) + defer n.waiters.Add(-1) + for { + n.mu.Lock() + ready := n.ready + n.mu.Unlock() + + // This check must stay after waiter registration and loading ready. It + // closes both windows in which notify could otherwise be lost. + if tryAcquire() { + return true + } + if span.stopped.Load() { + return false + } + select { + case <-ready: + case <-ctx.Done(): + return false + } + } +} + +func (n *eventMemoryNotifier) notify() { + // waiters is only a fast-path hint. A waiter that registers after this load + // rechecks memory and the span state before blocking, so observing a stale + // zero cannot lose a wakeup. Observing a stale nonzero only causes a harmless + // extra broadcast. + if n.waiters.Load() == 0 { + return + } + n.mu.Lock() + close(n.ready) + n.ready = make(chan struct{}) + n.mu.Unlock() +} + +// memoryQuotaController coordinates memory pressure from two sources: +// retained event memory and admitted initial scans. +// +// Event memory tracks bytes kept alive until downstream finishes consuming +// them. It may exceed the soft capacity temporarily, but the receive path +// blocks once it reaches the hard limit. +// +// Initial scans are charged by estimate instead of measured bytes. Each +// admitted scan starts from scanBaseSize, grows logarithmically with scan lag, +// and is capped at maxScanLagFactor times the base size. Scan admission +// compares max(event used, scan used) with the soft capacity: low-priority +// scans pause at pauseLowPriorityLimit and resume at +// resumeLowPriorityLimit, while high-priority scans continue to make progress. +type memoryQuotaController struct { + capacity uint64 + // used tracks event bytes retained until downstream finishes consuming them. + // Event accounting is on the receive hot path, so acquiring and releasing + // memory only use atomic operations while usage is below the hard limit. + used atomic.Uint64 + + // eventNotifier owns the wait protocol used after the hard limit is reached. + eventNotifier *eventMemoryNotifier + scanWaiters atomic.Int64 + + // scanMu guards scan admission state and scanReady. Scan admission happens + // once per region rather than once per event batch, so it is intentionally + // kept simple instead of adding atomics to every field. + scanMu sync.Mutex + // scanUsed tracks the estimated memory of all admitted initial scans. + scanUsed uint64 + level admissionLevel + // scanReady is replaced and closed when a memory transition can make a + // rejected scan eligible. Workers wait on this channel directly, avoiding a + // synchronous broadcast to every store and request worker. + scanReady chan struct{} + + pauseLowPriorityLimit uint64 + resumeLowPriorityLimit uint64 + hardLimit uint64 + + scanEstimate uint64 +} + +func newMemoryQuotaController(capacity, scanBaseSize uint64) *memoryQuotaController { + c := &memoryQuotaController{ + capacity: capacity, + level: admissionNormal, + pauseLowPriorityLimit: uint64(math.Ceil(float64(capacity) * defaultPauseLowPriorityRatio)), + resumeLowPriorityLimit: uint64(float64(capacity) * defaultResumeLowPriorityRatio), + hardLimit: uint64(float64(capacity) * defaultHardLimitRatio), + scanEstimate: scanBaseSize, + eventNotifier: newEventMemoryNotifier(), + scanReady: make(chan struct{}), + } + return c +} + +// WakeAll wakes quota waiters so they can observe cancellation or a stopped span. +func (c *memoryQuotaController) WakeAll() { + c.eventNotifier.notify() + c.NotifyScanAdmission() +} + +// AcquireScan admits one region scan and returns its memory estimate. +func (c *memoryQuotaController) AcquireScan( + region regionInfo, + currentTs uint64, +) (bytes uint64, retry <-chan struct{}, admitted bool) { + span := region.subscribedSpan + if span.stopped.Load() { + // Let stale tasks reach the worker's stopped-subscription cleanup path + // without consuming scan quota. + return 0, nil, true + } + + c.scanMu.Lock() + defer c.scanMu.Unlock() + c.refreshLevelLocked() + lowPriority := isLowPriorityScan(region, currentTs) + // Admission is based on the pressure before accounting this scan. This lets + // one scan make progress even when its estimate alone exceeds the threshold. + if lowPriority && c.level == admissionPauseLowPriority { + return 0, c.scanReady, false + } + bytes = c.estimateScanSizeLocked(region, currentTs) + c.scanUsed += bytes + c.refreshLevelLocked() + return bytes, nil, true +} + +// ReleaseScan releases the estimate owned by an admitted region scan. +func (c *memoryQuotaController) ReleaseScan(bytes uint64) { + if bytes == 0 { + return + } + c.scanMu.Lock() + previousLevel := c.level + c.scanUsed = subtractFloor(c.scanUsed, bytes) + c.refreshLevelLocked() + if c.level < previousLevel { + c.notifyScanAdmissionLocked() + } + c.scanMu.Unlock() +} + +// AcquireEvent accounts one event batch. Below the hard limit its hot path is +// a context check and an atomic compare-and-swap; it does not allocate or take +// a mutex. +func (c *memoryQuotaController) AcquireEvent( + ctx context.Context, + span *subscribedSpan, + bytes uint64, +) bool { + if ctx.Err() != nil { + return false + } + if c.tryAcquireEvent(bytes) { + return true + } + + start := time.Now() + acquired := c.eventNotifier.wait(ctx, span, func() bool { + return c.tryAcquireEvent(bytes) + }) + metrics.LogPullerMemoryQuotaEventWaitDuration.Observe(time.Since(start).Seconds()) + return acquired +} + +func (c *memoryQuotaController) tryAcquireEvent(bytes uint64) bool { + for { + used := c.used.Load() + if used > 0 && wouldExceed(used, bytes, c.hardLimit) { + return false + } + if bytes > math.MaxUint64-used { + return false + } + if c.used.CompareAndSwap(used, used+bytes) { + return true + } + } +} + +// ReleaseEvent releases event memory after downstream has consumed the event. +func (c *memoryQuotaController) ReleaseEvent(bytes uint64) { + if bytes == 0 { + return + } + used := c.used.Add(^(bytes - 1)) + previousUsed := used + bytes + if crossesDown(previousUsed, used, c.resumeLowPriorityLimit) { + c.refreshAdmissionAndNotify() + } + c.eventNotifier.notify() +} + +// NotifyScanAdmission wakes workers so they can recheck span state and admission. +func (c *memoryQuotaController) NotifyScanAdmission() { + c.scanMu.Lock() + c.notifyScanAdmissionLocked() + c.scanMu.Unlock() +} + +// UpdateMetrics reports the current event-memory and scan-admission state. +func (c *memoryQuotaController) UpdateMetrics() { + c.scanMu.Lock() + used := c.used.Load() + scanUsed := c.scanUsed + c.scanMu.Unlock() + + metrics.LogPullerMemoryQuota.WithLabelValues("max").Set(float64(c.capacity)) + metrics.LogPullerMemoryQuota.WithLabelValues("used").Set(float64(used)) + metrics.LogPullerMemoryQuota.WithLabelValues("scan_used").Set(float64(scanUsed)) + metrics.LogPullerMemoryQuotaEventWaiterCount.Set( + float64(c.eventNotifier.waiters.Load())) + metrics.LogPullerMemoryQuotaScanWaiterCount.Set( + float64(c.scanWaiters.Load())) +} + +func (c *memoryQuotaController) notifyScanAdmissionLocked() { + close(c.scanReady) + c.scanReady = make(chan struct{}) +} + +func (c *memoryQuotaController) refreshAdmissionAndNotify() { + c.scanMu.Lock() + previousLevel := c.level + c.refreshLevelLocked() + if c.level < previousLevel { + c.notifyScanAdmissionLocked() + } + c.scanMu.Unlock() +} + +func (c *memoryQuotaController) estimateScanSizeLocked(region regionInfo, currentTs uint64) uint64 { + raw := float64(c.scanEstimate) * scanLagFactor(region.resolvedTs(), currentTs) + estimate := uint64(raw) + if estimate < c.scanEstimate { + estimate = c.scanEstimate + } + maxEstimate := uint64(math.MaxUint64) + if c.scanEstimate <= math.MaxUint64/defaultMaxScanLagFactor { + maxEstimate = c.scanEstimate * defaultMaxScanLagFactor + } + if estimate > maxEstimate { + estimate = maxEstimate + } + if estimate == 0 { + estimate = c.scanEstimate + } + return estimate +} + +func scanLagFactor(startTs, currentTs uint64) float64 { + lag := regionScanLag(currentTs, startTs) + if lag <= 0 { + return 1 + } + return min(defaultMaxScanLagFactor, + 1+defaultScanLagWeight*math.Log2(1+float64(lag)/float64(defaultScanLagUnit))) +} + +func regionScanLag(currentTs, checkpointTs uint64) time.Duration { + currentTime := oracle.GetTimeFromTS(currentTs) + checkpointTime := oracle.GetTimeFromTS(checkpointTs) + if !currentTime.After(checkpointTime) { + return 0 + } + return currentTime.Sub(checkpointTime) +} + +func isLowPriorityScan(region regionInfo, _ uint64) bool { + return !isHighScanPriority(region.scanPriority) +} + +func (c *memoryQuotaController) refreshLevelLocked() { + // scanUsed predicts the event memory an initial scan may produce, so adding + // it to actual event bytes would count the same pressure twice. + pressure := max(c.used.Load(), c.scanUsed) + switch c.level { + case admissionPauseLowPriority: + if pressure <= c.resumeLowPriorityLimit { + c.level = admissionNormal + } + default: + if pressure >= c.pauseLowPriorityLimit { + c.level = admissionPauseLowPriority + } + } +} + +func wouldExceed(used, bytes, limit uint64) bool { + return bytes > limit || used > limit-bytes +} + +func crossesDown(previous, current, threshold uint64) bool { + return previous > threshold && current <= threshold +} + +func subtractFloor(value, delta uint64) uint64 { + if value < delta { + return 0 + } + return value - delta +} diff --git a/logservice/logpuller/memory_quota_test.go b/logservice/logpuller/memory_quota_test.go new file mode 100644 index 0000000000..41525f7b1e --- /dev/null +++ b/logservice/logpuller/memory_quota_test.go @@ -0,0 +1,441 @@ +// Copyright 2026 PingCAP, Inc. +// +// 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 logpuller + +import ( + "context" + "testing" + "time" + + "github.com/pingcap/kvproto/pkg/cdcpb" + "github.com/pingcap/ticdc/logservice/logpuller/regionlock" + "github.com/pingcap/ticdc/pkg/metrics" + "github.com/pingcap/ticdc/pkg/pdutil" + "github.com/prometheus/client_golang/prometheus/testutil" + "github.com/stretchr/testify/require" + "github.com/tikv/client-go/v2/oracle" +) + +type memoryQuotaTestState struct { + used uint64 + scanUsed uint64 + level admissionLevel +} + +func getMemoryQuotaTestState(quota *memoryQuotaController) memoryQuotaTestState { + quota.scanMu.Lock() + defer quota.scanMu.Unlock() + return memoryQuotaTestState{ + used: quota.used.Load(), + scanUsed: quota.scanUsed, + level: quota.level, + } +} + +func newTestQuotaSpan(subID SubscriptionID) *subscribedSpan { + span := &subscribedSpan{subID: subID} + span.resolvedTs.Store(oracle.GoTimeToTS(time.Now())) + return span +} + +func newTestQuotaRegion(span *subscribedSpan) regionInfo { + state := ®ionlock.LockedRangeState{} + state.ResolvedTs.Store(span.resolvedTs.Load()) + return regionInfo{ + subscribedSpan: span, + lockedRangeState: state, + } +} + +func newTestQuotaRegionWithPriority( + span *subscribedSpan, + priority cdcpb.ScanPriority, +) regionInfo { + region := newTestQuotaRegion(span) + region.scanPriority = priority + return region +} + +func setTestQuotaSpanLag(span *subscribedSpan, lag time.Duration) uint64 { + now := time.Now() + span.resolvedTs.Store(oracle.GoTimeToTS(now.Add(-lag))) + return oracle.GoTimeToTS(now) +} + +func TestMemoryQuotaUpdateMetrics(t *testing.T) { + quota := newMemoryQuotaController(66, 8) + span := newTestQuotaSpan(1) + require.True(t, quota.AcquireEvent(context.Background(), span, 55)) + t.Cleanup(func() { quota.ReleaseEvent(55) }) + + quota.scanMu.Lock() + quota.scanUsed = 7 + quota.scanMu.Unlock() + quota.eventNotifier.waiters.Store(2) + t.Cleanup(func() { quota.eventNotifier.waiters.Store(0) }) + quota.scanWaiters.Store(3) + t.Cleanup(func() { quota.scanWaiters.Store(0) }) + + quota.UpdateMetrics() + + require.Equal(t, float64(66), testutil.ToFloat64( + metrics.LogPullerMemoryQuota.WithLabelValues("max"))) + require.Equal(t, float64(55), testutil.ToFloat64( + metrics.LogPullerMemoryQuota.WithLabelValues("used"))) + require.Equal(t, float64(7), testutil.ToFloat64( + metrics.LogPullerMemoryQuota.WithLabelValues("scan_used"))) + require.Equal(t, float64(2), + testutil.ToFloat64(metrics.LogPullerMemoryQuotaEventWaiterCount)) + require.Equal(t, float64(3), + testutil.ToFloat64(metrics.LogPullerMemoryQuotaScanWaiterCount)) +} + +func TestMemoryQuotaAdmissionLevels(t *testing.T) { + quota := newMemoryQuotaController(100, 10) + lowPrioritySpan := newTestQuotaSpan(1) + highPrioritySpan := newTestQuotaSpan(2) + lowPriorityTs := setTestQuotaSpanLag(lowPrioritySpan, time.Hour) + highPriorityTs := setTestQuotaSpanLag(highPrioritySpan, time.Hour) + + require.True(t, quota.AcquireEvent(context.Background(), highPrioritySpan, 5)) + require.True(t, quota.AcquireEvent(context.Background(), highPrioritySpan, 10)) + _, _, admitted := quota.AcquireScan( + newTestQuotaRegionWithPriority(lowPrioritySpan, cdcpb.ScanPriority_SCAN_PRIORITY_LOW), + lowPriorityTs, + ) + require.False(t, admitted) + + scanBytes, _, admitted := quota.AcquireScan( + newTestQuotaRegionWithPriority(highPrioritySpan, cdcpb.ScanPriority_SCAN_PRIORITY_HIGH), + highPriorityTs, + ) + require.True(t, admitted) + quota.ReleaseScan(scanBytes) + + require.True(t, quota.AcquireEvent(context.Background(), highPrioritySpan, 45)) + require.True(t, quota.AcquireEvent(context.Background(), highPrioritySpan, 20)) + scanBytes, _, admitted = quota.AcquireScan( + newTestQuotaRegionWithPriority(highPrioritySpan, cdcpb.ScanPriority_SCAN_PRIORITY_HIGH), + highPriorityTs, + ) + require.True(t, admitted) + quota.ReleaseScan(scanBytes) + + quota.ReleaseEvent(20) + state := getMemoryQuotaTestState(quota) + require.Equal(t, admissionPauseLowPriority, state.level) + quota.ReleaseEvent(45) + state = getMemoryQuotaTestState(quota) + require.Equal(t, admissionPauseLowPriority, state.level) + quota.ReleaseEvent(10) + state = getMemoryQuotaTestState(quota) + require.Equal(t, admissionNormal, state.level) + quota.ReleaseEvent(5) +} + +func TestMemoryQuotaSpanStopKeepsOwnedMemoryUntilRelease(t *testing.T) { + quota := newMemoryQuotaController(100, 10) + span1 := newTestQuotaSpan(1) + span2 := newTestQuotaSpan(2) + + require.True(t, quota.AcquireEvent(context.Background(), span1, 30)) + require.True(t, quota.AcquireEvent(context.Background(), span2, 40)) + scanBytes, _, admitted := quota.AcquireScan( + newTestQuotaRegionWithPriority(span1, cdcpb.ScanPriority_SCAN_PRIORITY_HIGH), + span1.resolvedTs.Load(), + ) + require.True(t, admitted) + require.NotZero(t, scanBytes) + + span1.stopped.Store(true) + quota.WakeAll() + state := getMemoryQuotaTestState(quota) + require.Equal(t, uint64(70), state.used) + require.Equal(t, scanBytes, state.scanUsed) + + quota.ReleaseEvent(30) + quota.ReleaseScan(scanBytes) + state = getMemoryQuotaTestState(quota) + require.Equal(t, uint64(40), state.used) + + // Late tasks reach the stopped-subscription cleanup path without consuming + // scan quota. + scanBytes, _, admitted = quota.AcquireScan( + newTestQuotaRegionWithPriority(span1, cdcpb.ScanPriority_SCAN_PRIORITY_HIGH), + span1.resolvedTs.Load(), + ) + require.True(t, admitted) + require.Zero(t, scanBytes) + + quota.ReleaseEvent(40) + state = getMemoryQuotaTestState(quota) + require.Zero(t, state.used) +} + +func TestMemoryQuotaBlockedEventStopsWhenSpanStops(t *testing.T) { + quota := newMemoryQuotaController(100, 10) + quota.hardLimit = 100 + span := newTestQuotaSpan(1) + + require.True(t, quota.AcquireEvent(context.Background(), span, 100)) + acquired := make(chan bool, 1) + go func() { + acquired <- quota.AcquireEvent(context.Background(), span, 1) + }() + + select { + case <-acquired: + t.Fatal("event memory should wait at the hard limit") + case <-time.After(100 * time.Millisecond): + } + + span.stopped.Store(true) + quota.WakeAll() + select { + case ok := <-acquired: + require.False(t, ok) + case <-time.After(time.Second): + t.Fatal("stopping the subscription did not wake the blocked event") + } + quota.ReleaseEvent(100) +} + +func TestMemoryQuotaBlockedEventResumesAfterRelease(t *testing.T) { + quota := newMemoryQuotaController(100, 10) + span := newTestQuotaSpan(1) + + require.True(t, quota.AcquireEvent(context.Background(), span, 200)) + acquired := make(chan bool, 1) + go func() { + acquired <- quota.AcquireEvent(context.Background(), span, 1) + }() + + select { + case <-acquired: + t.Fatal("event memory should wait at the hard limit") + case <-time.After(100 * time.Millisecond): + } + + quota.ReleaseEvent(200) + select { + case ok := <-acquired: + require.True(t, ok) + quota.ReleaseEvent(1) + case <-time.After(time.Second): + t.Fatal("event memory did not resume after memory was released") + } +} + +func TestMemoryQuotaBlockedEventStopsOnContextCancellation(t *testing.T) { + quota := newMemoryQuotaController(100, 10) + quota.hardLimit = 100 + span := newTestQuotaSpan(1) + ctx, cancel := context.WithCancel(context.Background()) + + require.True(t, quota.AcquireEvent(context.Background(), span, 100)) + acquired := make(chan bool, 1) + go func() { + acquired <- quota.AcquireEvent(ctx, span, 1) + }() + require.Eventually(t, func() bool { + return quota.eventNotifier.waiters.Load() == 1 + }, time.Second, time.Millisecond) + + // Cancellation must stop the waiter without a memory release or an explicit + // quota notification. + cancel() + select { + case ok := <-acquired: + require.False(t, ok) + case <-time.After(time.Second): + t.Fatal("context cancellation did not stop the blocked event") + } + quota.ReleaseEvent(100) +} + +func TestMemoryQuotaConcurrentWaitersDoNotLoseWakeups(t *testing.T) { + const waiterCount = 32 + quota := newMemoryQuotaController(100, 10) + quota.hardLimit = 1 + span := newTestQuotaSpan(1) + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + // Hold the only available byte until every goroutine is waiting. Releasing + // it wakes all waiters; each successful waiter then releases it for the next. + require.True(t, quota.AcquireEvent(ctx, span, 1)) + results := make(chan bool, waiterCount) + for range waiterCount { + go func() { + acquired := quota.AcquireEvent(ctx, span, 1) + if acquired { + quota.ReleaseEvent(1) + } + results <- acquired + }() + } + require.Eventually(t, func() bool { + return quota.eventNotifier.waiters.Load() == waiterCount + }, time.Second, time.Millisecond) + + quota.ReleaseEvent(1) + for range waiterCount { + select { + case acquired := <-results: + require.True(t, acquired) + case <-ctx.Done(): + t.Fatal("event waiter did not make progress") + } + } + state := getMemoryQuotaTestState(quota) + require.Zero(t, state.used) +} + +func TestMemoryQuotaLowPriorityScanUsesCurrentPressure(t *testing.T) { + quota := newMemoryQuotaController(100, 20) + span := newTestQuotaSpan(1) + currentTs := setTestQuotaSpanLag(span, time.Hour) + region := newTestQuotaRegionWithPriority(span, cdcpb.ScanPriority_SCAN_PRIORITY_LOW) + + bytes1, _, admitted := quota.AcquireScan(region, currentTs) + require.True(t, admitted) + require.NotZero(t, bytes1) + state := getMemoryQuotaTestState(quota) + require.Greater(t, state.scanUsed, quota.pauseLowPriorityLimit) + + _, _, admitted = quota.AcquireScan(region, currentTs) + require.False(t, admitted) + + quota.ReleaseScan(bytes1) + bytes2, _, admitted := quota.AcquireScan(region, currentTs) + require.True(t, admitted) + quota.ReleaseScan(bytes2) +} + +func TestMemoryQuotaLowLagScanBypassesWarmingGate(t *testing.T) { + quota := newMemoryQuotaController(100, 10) + span := newTestQuotaSpan(1) + currentTs := setTestQuotaSpanLag(span, time.Minute) + + require.True(t, quota.AcquireEvent(context.Background(), span, 20)) + scanBytes, _, admitted := quota.AcquireScan( + newTestQuotaRegionWithPriority(span, cdcpb.ScanPriority_SCAN_PRIORITY_HIGH), + currentTs, + ) + require.True(t, admitted) + require.NotZero(t, scanBytes) + state := getMemoryQuotaTestState(quota) + require.NotZero(t, state.scanUsed) + + quota.ReleaseScan(scanBytes) + quota.ReleaseEvent(20) +} + +func TestAdmissionWaitsForMemoryAndReleasesScanMemory(t *testing.T) { + quota := newMemoryQuotaController(100, 10) + span := newTestQuotaSpan(1) + currentTs := setTestQuotaSpanLag(span, time.Hour) + clock := pdutil.NewClock4Test().(*pdutil.Clock4Test) + clock.SetTS(currentTs) + controller := newRegionAdmissionController(1, 1, quota, clock) + + require.True(t, quota.AcquireEvent(context.Background(), span, 20)) + region := newTestQuotaRegionWithPriority(span, cdcpb.ScanPriority_SCAN_PRIORITY_LOW) + require.True(t, controller.submit(newRegionPriorityTask(region, 1))) + + type popResult struct { + req *regionReq + err error + } + result := make(chan popResult, 1) + go func() { + req, err := controller.pop(context.Background(), nil) + result <- popResult{req: req, err: err} + }() + select { + case <-result: + t.Fatal("low-priority scan should wait while memory is under pressure") + case <-time.After(100 * time.Millisecond): + } + + quota.ReleaseEvent(20) + var resultValue popResult + select { + case resultValue = <-result: + case <-time.After(time.Second): + t.Fatal("scan admission was not notified after memory became available") + } + require.NoError(t, resultValue.err) + req := resultValue.req + state := getMemoryQuotaTestState(quota) + require.NotZero(t, state.scanUsed) + require.True(t, req.abort()) + state = getMemoryQuotaTestState(quota) + require.Zero(t, state.scanUsed) +} + +func TestAdmissionWakesWhenBlockedSpanStops(t *testing.T) { + quota := newMemoryQuotaController(100, 10) + span := newTestQuotaSpan(1) + currentTs := setTestQuotaSpanLag(span, time.Hour) + clock := pdutil.NewClock4Test().(*pdutil.Clock4Test) + clock.SetTS(currentTs) + controller := newRegionAdmissionController(1, 1, quota, clock) + + require.True(t, quota.AcquireEvent(context.Background(), span, 20)) + require.True(t, controller.submit(newRegionPriorityTask( + newTestQuotaRegionWithPriority(span, cdcpb.ScanPriority_SCAN_PRIORITY_LOW), 1))) + + type popResult struct { + req *regionReq + err error + } + result := make(chan popResult, 1) + go func() { + req, err := controller.pop(context.Background(), nil) + result <- popResult{req: req, err: err} + }() + select { + case <-result: + t.Fatal("low-priority scan should wait while memory is under pressure") + case <-time.After(100 * time.Millisecond): + } + + span.stopped.Store(true) + quota.WakeAll() + select { + case result := <-result: + require.NoError(t, result.err) + require.Zero(t, result.req.scanBytes) + require.True(t, result.req.abort()) + case <-time.After(time.Second): + t.Fatal("stopping the span did not wake scan admission") + } + quota.ReleaseEvent(20) +} + +func BenchmarkMemoryQuotaEventAccounting(b *testing.B) { + quota := newMemoryQuotaController(1024*1024*1024, 8*1024*1024) + span := newTestQuotaSpan(1) + ctx := context.Background() + b.ReportAllocs() + b.ResetTimer() + for b.Loop() { + if !quota.AcquireEvent(ctx, span, 1) { + b.Fatal("failed to acquire event memory") + } + quota.ReleaseEvent(1) + } +} diff --git a/logservice/logpuller/mock_upstream.go b/logservice/logpuller/mock_upstream_test.go similarity index 100% rename from logservice/logpuller/mock_upstream.go rename to logservice/logpuller/mock_upstream_test.go diff --git a/logservice/logpuller/priority_task.go b/logservice/logpuller/priority_task.go index 56bc7f72cc..38cd813c8b 100644 --- a/logservice/logpuller/priority_task.go +++ b/logservice/logpuller/priority_task.go @@ -13,130 +13,41 @@ package logpuller -import ( - "time" +import "github.com/pingcap/kvproto/pkg/cdcpb" - "github.com/pingcap/kvproto/pkg/cdcpb" - "github.com/tikv/client-go/v2/oracle" -) - -// TaskType represents the type of region task -type TaskType int - -const ( - // TaskHighPrior represents region error or region change - // This type has the highest priority - TaskHighPrior TaskType = iota - // TaskLowPrior represents new subscription - // This type has the lowest priority - TaskLowPrior -) - -const ( - highPriorityBase = 0 - lowPriorityBase = 60 * 60 * 24 // 1 day - forcedPriorityBase = 60 * 60 // 60 minutes -) - -func (t TaskType) String() string { - switch t { - case TaskHighPrior: - return "high" - case TaskLowPrior: - return "low" - default: - return "unknown" - } -} - -func (t TaskType) scanPriority() cdcpb.ScanPriority { - switch t { - case TaskHighPrior: - return cdcpb.ScanPriority_SCAN_PRIORITY_HIGH - case TaskLowPrior: - return cdcpb.ScanPriority_SCAN_PRIORITY_LOW - default: - return cdcpb.ScanPriority_SCAN_PRIORITY_LOW - } -} - -func taskTypeFromScanPriority(priority cdcpb.ScanPriority) TaskType { +func normalizeScanPriority(priority cdcpb.ScanPriority) cdcpb.ScanPriority { if priority == cdcpb.ScanPriority_SCAN_PRIORITY_HIGH { - return TaskHighPrior + return cdcpb.ScanPriority_SCAN_PRIORITY_HIGH } - return TaskLowPrior -} - -func normalizeScanPriority(priority cdcpb.ScanPriority) cdcpb.ScanPriority { - return taskTypeFromScanPriority(priority).scanPriority() + return cdcpb.ScanPriority_SCAN_PRIORITY_LOW } -// PriorityTask is the interface for priority-based tasks -// It implements heap.Item interface -type PriorityTask interface { - // Priority returns the priority value, lower value means higher priority - Priority() int - - // GetRegionInfo returns the underlying regionInfo - GetRegionInfo() regionInfo - - // heap.Item interface methods - SetHeapIndex(int) - GetHeapIndex() int - LessThan(PriorityTask) bool +func isHighScanPriority(priority cdcpb.ScanPriority) bool { + return normalizeScanPriority(priority) == cdcpb.ScanPriority_SCAN_PRIORITY_HIGH } -// regionPriorityTask implements PriorityTask interface type regionPriorityTask struct { - taskType TaskType - createTime time.Time regionInfo regionInfo + sequence uint64 heapIndex int // for heap.Item interface - currentTs uint64 } -// NewRegionPriorityTask creates a new priority task for region -func NewRegionPriorityTask(taskType TaskType, regionInfo regionInfo, currentTs uint64) PriorityTask { +// newRegionPriorityTask creates a new priority task for region. +func newRegionPriorityTask(regionInfo regionInfo, sequence uint64) *regionPriorityTask { + regionInfo.scanPriority = normalizeScanPriority(regionInfo.scanPriority) return ®ionPriorityTask{ - taskType: taskType, - createTime: time.Now(), regionInfo: regionInfo, + sequence: sequence, heapIndex: 0, // 0 means not in heap - currentTs: currentTs, } } -// Priority calculates the priority based on task type and wait time -// Lower value means higher priority -func (pt *regionPriorityTask) Priority() int { - // Base priority based on task type - basePriority := 0 - switch pt.taskType { - case TaskHighPrior: - basePriority = highPriorityBase // Highest priority - case TaskLowPrior: - basePriority = lowPriorityBase // Lowest priority - } - - // Add time-based priority bonus - // Wait time in seconds, longer wait time means higher priority (lower value) - waitTime := time.Since(pt.createTime) - timeBonus := int(waitTime.Seconds()) - - // ResolvedTsLag in seconds, longer lag means lower priority (higher value) - resolvedTsLag := oracle.GetTimeFromTS(pt.currentTs).Sub(oracle.GetTimeFromTS(pt.regionInfo.subscribedSpan.resolvedTs.Load())) - resolvedTsLagPenalty := int(resolvedTsLag.Seconds()) - - priority := basePriority - timeBonus + resolvedTsLagPenalty - if priority < 0 { - priority = 0 - } - return priority +func (pt *regionPriorityTask) priority() cdcpb.ScanPriority { + return normalizeScanPriority(pt.regionInfo.scanPriority) } -// GetRegionInfo returns the underlying regionInfo -func (pt *regionPriorityTask) GetRegionInfo() regionInfo { - return pt.regionInfo +func (pt *regionPriorityTask) canUseMaxWindow() bool { + return isHighScanPriority(pt.regionInfo.scanPriority) } // SetHeapIndex sets the heap index for heap.Item interface @@ -149,8 +60,11 @@ func (pt *regionPriorityTask) GetHeapIndex() int { return pt.heapIndex } -// LessThan implements heap.Item interface -// Returns true if this task has higher priority (lower priority value) than the other task -func (pt *regionPriorityTask) LessThan(other PriorityTask) bool { - return pt.Priority() < other.Priority() +// LessThan implements heap.Item interface. Tasks in the same priority class are +// processed in submission order. +func (pt *regionPriorityTask) LessThan(other *regionPriorityTask) bool { + if isHighScanPriority(pt.regionInfo.scanPriority) != isHighScanPriority(other.regionInfo.scanPriority) { + return isHighScanPriority(pt.regionInfo.scanPriority) + } + return pt.sequence < other.sequence } diff --git a/logservice/logpuller/priority_task_test.go b/logservice/logpuller/priority_task_test.go index 921f227102..41501181fb 100644 --- a/logservice/logpuller/priority_task_test.go +++ b/logservice/logpuller/priority_task_test.go @@ -8,249 +8,128 @@ // // 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 logpuller import ( - "sync/atomic" "testing" "time" "github.com/pingcap/kvproto/pkg/cdcpb" "github.com/pingcap/ticdc/heartbeatpb" + "github.com/pingcap/ticdc/logservice/logpuller/regionlock" "github.com/pingcap/ticdc/utils/priorityqueue" "github.com/stretchr/testify/require" "github.com/tikv/client-go/v2/oracle" "github.com/tikv/client-go/v2/tikv" ) -func TestTaskTypeScanPriorityMapping(t *testing.T) { - require.Equal(t, cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, TaskHighPrior.scanPriority()) - require.Equal(t, cdcpb.ScanPriority_SCAN_PRIORITY_LOW, TaskLowPrior.scanPriority()) - require.Equal(t, TaskHighPrior, taskTypeFromScanPriority(cdcpb.ScanPriority_SCAN_PRIORITY_HIGH)) - require.Equal(t, TaskLowPrior, taskTypeFromScanPriority(cdcpb.ScanPriority_SCAN_PRIORITY_LOW)) - require.Equal(t, TaskLowPrior, taskTypeFromScanPriority(cdcpb.ScanPriority_SCAN_PRIORITY_UNKNOWN)) - require.Equal(t, cdcpb.ScanPriority_SCAN_PRIORITY_LOW, normalizeScanPriority(cdcpb.ScanPriority_SCAN_PRIORITY_UNKNOWN)) -} - -// TestPriorityCalculationLogic tests the priority calculation logic in isolation -func TestPriorityCalculationLogic(t *testing.T) { - currentTime := time.Now() - currentTs := oracle.GoTimeToTS(currentTime) - - // Test cases for priority calculation - tests := []struct { - name string - taskType TaskType - resolvedTsOffsetSeconds int64 // Offset relative to currentTs (negative means resolvedTs is older) - waitTimeSeconds int // Task wait time - description string - }{ - { - name: "high_priority_new_resolvedTs", - taskType: TaskHighPrior, - resolvedTsOffsetSeconds: -5, // resolvedTs is 5 seconds earlier than currentTs - waitTimeSeconds: 10, // Waited for 10 seconds - description: "High priority task with newer resolvedTs", - }, - { - name: "high_priority_old_resolvedTs", - taskType: TaskHighPrior, - resolvedTsOffsetSeconds: -30, // resolvedTs is 30 seconds earlier than currentTs - waitTimeSeconds: 10, // Waited for 10 seconds - description: "High priority task with older resolvedTs", - }, - { - name: "low_priority_new_resolvedTs", - taskType: TaskLowPrior, - resolvedTsOffsetSeconds: -5, // resolvedTs is 5 seconds earlier than currentTs - waitTimeSeconds: 10, // Waited for 10 seconds - description: "Low priority task with newer resolvedTs", - }, - { - name: "low_priority_old_resolvedTs", - taskType: TaskLowPrior, - resolvedTsOffsetSeconds: -30, // resolvedTs is 30 seconds earlier than currentTs - waitTimeSeconds: 10, // Waited for 10 seconds - description: "Low priority task with older resolvedTs", - }, - } - - var priorities []int - var taskDescriptions []string - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - // Calculate resolvedTs: currentTs + offset - resolvedTime := oracle.GetTimeFromTS(currentTs).Add(time.Duration(tt.resolvedTsOffsetSeconds) * time.Second) - resolvedTs := oracle.GoTimeToTS(resolvedTime) - - // Simulate priority calculation logic - priority := calculatePriorityDirectly(tt.taskType, currentTs, resolvedTs, tt.waitTimeSeconds) - - t.Logf("%s: Priority = %d", tt.description, priority) - priorities = append(priorities, priority) - taskDescriptions = append(taskDescriptions, tt.description) - }) +func newPriorityTestRegion( + regionID uint64, + checkpointTs uint64, +) regionInfo { + span := heartbeatpb.TableSpan{TableID: 1, StartKey: []byte("a"), EndKey: []byte("z")} + state := ®ionlock.LockedRangeState{} + state.ResolvedTs.Store(checkpointTs) + return regionInfo{ + verID: tikv.NewRegionVerID(regionID, 1, 1), + span: span, + subscribedSpan: &subscribedSpan{subID: 1, startTs: checkpointTs, span: span}, + lockedRangeState: state, } - - // Verify priority order - t.Run("verify_priority_order", func(t *testing.T) { - require.Equal(t, 4, len(priorities)) - - highPriorNewResolvedTs := priorities[0] // High priority with new resolvedTs - highPriorOldResolvedTs := priorities[1] // High priority with old resolvedTs - lowPriorNewResolvedTs := priorities[2] // Low priority with new resolvedTs - lowPriorOldResolvedTs := priorities[3] // Low priority with old resolvedTs - - t.Logf("Priority comparison:") - for i, desc := range taskDescriptions { - t.Logf(" %s: %d", desc, priorities[i]) - } - - // Core verification: For the same task type, newer resolvedTs (smaller lag) should have higher priority (smaller value) - require.Less(t, highPriorNewResolvedTs, highPriorOldResolvedTs, - "For the same task type, tasks with newer resolvedTs should have higher priority") - require.Less(t, lowPriorNewResolvedTs, lowPriorOldResolvedTs, - "For the same task type, tasks with newer resolvedTs should have higher priority") - - // Verify: High priority tasks always have higher priority than low priority tasks - require.Less(t, highPriorNewResolvedTs, lowPriorNewResolvedTs, - "High priority tasks should have higher priority than low priority tasks") - require.Less(t, highPriorOldResolvedTs, lowPriorOldResolvedTs, - "Even with older resolvedTs, high priority tasks should still have higher priority than low priority tasks") - }) } -// calculatePriorityDirectly directly calculates priority for testing -// Copies the logic from regionPriorityTask.Priority() -func calculatePriorityDirectly(taskType TaskType, currentTs, resolvedTs uint64, waitTimeSeconds int) int { - // Base priority based on task type - basePriority := 0 - switch taskType { - case TaskHighPrior: - basePriority = highPriorityBase // 1200 - case TaskLowPrior: - basePriority = lowPriorityBase // 3600 - } - - // Add time-based priority bonus - // Wait time in seconds, longer wait time means higher priority (lower value) - timeBonus := waitTimeSeconds - - // Calculate resolvedTs lag - resolvedTsLag := oracle.GetTimeFromTS(currentTs).Sub(oracle.GetTimeFromTS(resolvedTs)) - resolvedTsLagBonus := int(resolvedTsLag.Seconds()) - - priority := basePriority - timeBonus + resolvedTsLagBonus +func withScanPriority(region regionInfo, priority cdcpb.ScanPriority) regionInfo { + region.scanPriority = priority + return region +} - if priority < 0 { - priority = 0 - } - return priority +func TestNormalizeScanPriority(t *testing.T) { + require.Equal(t, cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, normalizeScanPriority(cdcpb.ScanPriority_SCAN_PRIORITY_HIGH)) + require.Equal(t, cdcpb.ScanPriority_SCAN_PRIORITY_LOW, normalizeScanPriority(cdcpb.ScanPriority_SCAN_PRIORITY_LOW)) + require.Equal(t, cdcpb.ScanPriority_SCAN_PRIORITY_LOW, normalizeScanPriority(cdcpb.ScanPriority_SCAN_PRIORITY_UNKNOWN)) + require.True(t, isHighScanPriority(cdcpb.ScanPriority_SCAN_PRIORITY_HIGH)) + require.False(t, isHighScanPriority(cdcpb.ScanPriority_SCAN_PRIORITY_LOW)) + require.False(t, isHighScanPriority(cdcpb.ScanPriority_SCAN_PRIORITY_UNKNOWN)) } -func TestResolvedTsLagLogic(t *testing.T) { +func TestRegionPriorityTaskQueueOrder(t *testing.T) { + queue := priorityqueue.New[*regionPriorityTask]() currentTime := time.Now() - currentTs := oracle.GoTimeToTS(currentTime) - - t.Run("test_resolvedTs_lag_calculation_logic", func(t *testing.T) { - // Scenario 1: resolvedTs is 10 seconds earlier than currentTs (resolvedTs is older) - resolvedTs1 := oracle.GoTimeToTS(currentTime.Add(-10 * time.Second)) - lag1 := oracle.GetTimeFromTS(currentTs).Sub(oracle.GetTimeFromTS(resolvedTs1)) - t.Logf("resolvedTs is 10 seconds earlier, lag = %v (%.0f seconds)", lag1, lag1.Seconds()) - - // Scenario 2: resolvedTs is 1 second earlier than currentTs (resolvedTs is newer) - resolvedTs2 := oracle.GoTimeToTS(currentTime.Add(-1 * time.Second)) - lag2 := oracle.GetTimeFromTS(currentTs).Sub(oracle.GetTimeFromTS(resolvedTs2)) - t.Logf("resolvedTs is 1 second earlier, lag = %v (%.0f seconds)", lag2, lag2.Seconds()) - // Verify: newer resolvedTs should have smaller lag - require.Less(t, lag2, lag1, "newer resolvedTs should have smaller lag") + lowTask := newRegionPriorityTask( + withScanPriority( + newPriorityTestRegion(1, oracle.GoTimeToTS(currentTime.Add(-time.Hour))), + cdcpb.ScanPriority_SCAN_PRIORITY_LOW, + ), + 3, + ) + highTask1 := newRegionPriorityTask( + withScanPriority( + newPriorityTestRegion(2, oracle.GoTimeToTS(currentTime.Add(-10*time.Minute))), + cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, + ), + 2, + ) + highTask2 := newRegionPriorityTask( + withScanPriority( + newPriorityTestRegion(3, oracle.GoTimeToTS(currentTime.Add(-time.Hour))), + cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, + ), + 1, + ) - // Calculate the impact on priority - priority1 := calculatePriorityDirectly(TaskHighPrior, currentTs, resolvedTs1, 5) - priority2 := calculatePriorityDirectly(TaskHighPrior, currentTs, resolvedTs2, 5) - - t.Logf("Priority with resolvedTs 10 seconds old: %d", priority1) - t.Logf("Priority with resolvedTs 1 second old: %d", priority2) + require.True(t, queue.Push(lowTask)) + require.True(t, queue.Push(highTask1)) + require.True(t, queue.Push(highTask2)) - // Verify: newer resolvedTs should have higher priority (smaller value) - require.Less(t, priority2, priority1, - "tasks with newer resolvedTs should have higher priority") - }) + for _, expectedRegionID := range []uint64{3, 2, 1} { + task, err := queue.Pop(t.Context()) + require.NoError(t, err) + require.Equal(t, expectedRegionID, task.regionInfo.verID.GetID()) + } } -func TestEdgeCases(t *testing.T) { +func TestRegionPriorityTaskFIFOWithinPriority(t *testing.T) { + queue := priorityqueue.New[*regionPriorityTask]() currentTime := time.Now() - currentTs := oracle.GoTimeToTS(currentTime) - - t.Run("resolvedTs in the future", func(t *testing.T) { - resolvedTs := oracle.GoTimeToTS(currentTime.Add(5 * time.Second)) - - lag := oracle.GetTimeFromTS(currentTs).Sub(oracle.GetTimeFromTS(resolvedTs)) - t.Logf("resolvedTs in the future 5 seconds, lag = %v (%.0f seconds)", lag, lag.Seconds()) - - priority := calculatePriorityDirectly(TaskHighPrior, currentTs, resolvedTs, 5) - t.Logf("resolvedTs in the future priority: %d", priority) - - require.GreaterOrEqual(t, priority, 0, "priority should not be less than 0") - }) + checkpointTs := oracle.GoTimeToTS(currentTime.Add(-time.Hour)) - t.Run("different wait time impact", func(t *testing.T) { - resolvedTs := oracle.GoTimeToTS(currentTime.Add(-10 * time.Second)) + first := newRegionPriorityTask( + withScanPriority(newPriorityTestRegion(1, checkpointTs), cdcpb.ScanPriority_SCAN_PRIORITY_HIGH), 1) + second := newRegionPriorityTask( + withScanPriority(newPriorityTestRegion(2, checkpointTs), cdcpb.ScanPriority_SCAN_PRIORITY_HIGH), 2) - priority1 := calculatePriorityDirectly(TaskHighPrior, currentTs, resolvedTs, 2) - priority2 := calculatePriorityDirectly(TaskHighPrior, currentTs, resolvedTs, 10) + require.True(t, queue.Push(second)) + require.True(t, queue.Push(first)) - t.Logf("wait 2 seconds priority: %d", priority1) - t.Logf("wait 10 seconds priority: %d", priority2) - - // wait time longer task priority should be higher - require.Less(t, priority2, priority1, "wait time longer task priority should be higher") - }) + task, err := queue.Pop(t.Context()) + require.NoError(t, err) + require.Equal(t, uint64(1), task.regionInfo.verID.GetID()) + task, err = queue.Pop(t.Context()) + require.NoError(t, err) + require.Equal(t, uint64(2), task.regionInfo.verID.GetID()) } -func TestRegionPriorityTaskQueueOrder(t *testing.T) { - queue := priorityqueue.New[PriorityTask]() - ctx := t.Context() - - currentTs := oracle.GoTimeToTS(time.Now()) - verID := tikv.NewRegionVerID(1, 1, 1) - span := heartbeatpb.TableSpan{TableID: 1, StartKey: []byte("a"), EndKey: []byte("z")} +func TestRegionPriorityTaskUsesHighPriorityWindow(t *testing.T) { + highTask := newRegionPriorityTask( + withScanPriority(newPriorityTestRegion(1, 1), cdcpb.ScanPriority_SCAN_PRIORITY_HIGH), 1) + lowTask := newRegionPriorityTask( + withScanPriority(newPriorityTestRegion(2, 1), cdcpb.ScanPriority_SCAN_PRIORITY_LOW), 2) - subscribedSpan := &subscribedSpan{ - resolvedTs: atomic.Uint64{}, - } - subscribedSpan.resolvedTs.Store(oracle.GoTimeToTS(time.Now().Add(-time.Second))) - - regionInfo := regionInfo{ - verID: verID, - span: span, - subscribedSpan: subscribedSpan, - } - - errorTask := NewRegionPriorityTask(TaskHighPrior, regionInfo, currentTs+1) - highTask := NewRegionPriorityTask(TaskHighPrior, regionInfo, currentTs) - lowTask := NewRegionPriorityTask(TaskLowPrior, regionInfo, currentTs) - - require.True(t, queue.Push(lowTask)) - require.True(t, queue.Push(errorTask)) - require.True(t, queue.Push(highTask)) - - first, err := queue.Pop(ctx) - require.NoError(t, err) - require.Equal(t, TaskHighPrior, first.(*regionPriorityTask).taskType) - - second, err := queue.Pop(ctx) - require.NoError(t, err) - require.Equal(t, TaskHighPrior, second.(*regionPriorityTask).taskType) + require.True(t, highTask.canUseMaxWindow()) + require.False(t, lowTask.canUseMaxWindow()) +} - third, err := queue.Pop(ctx) - require.NoError(t, err) - require.Equal(t, TaskLowPrior, third.(*regionPriorityTask).taskType) +func TestRegionPriorityTaskRefreshesRegionInfoBetweenStages(t *testing.T) { + region := withScanPriority(newPriorityTestRegion(1, 1), cdcpb.ScanPriority_SCAN_PRIORITY_LOW) + task := newRegionPriorityTask(region, 1) + require.Equal(t, cdcpb.ScanPriority_SCAN_PRIORITY_LOW, task.priority()) - require.Equal(t, 0, queue.Len()) + region.scanPriority = cdcpb.ScanPriority_SCAN_PRIORITY_HIGH + task.regionInfo = region + require.Equal(t, cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, task.priority()) } diff --git a/logservice/logpuller/region_admission_controller.go b/logservice/logpuller/region_admission_controller.go new file mode 100644 index 0000000000..7ebf3d1ff1 --- /dev/null +++ b/logservice/logpuller/region_admission_controller.go @@ -0,0 +1,289 @@ +// Copyright 2026 PingCAP, Inc. +// +// 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, +// See the License for the specific language governing permissions and +// limitations under the License. + +package logpuller + +import ( + "context" + "math" + "sync" + "sync/atomic" + "time" + + "github.com/pingcap/log" + "github.com/pingcap/ticdc/pkg/metrics" + "github.com/pingcap/ticdc/pkg/pdutil" + "github.com/pingcap/ticdc/utils/heap" + "go.uber.org/zap" +) + +const ( + abnormalRequestDurationInSec = 60 * 60 * 2 // 2 hours +) + +// regionReq is an admission lease for one sent-but-not-initialized region. +// finish and abort are idempotent and return the lease to its worker controller. +type regionReq struct { + regionInfo regionInfo + createTime time.Time + controller *regionAdmissionController + scanBytes uint64 + released atomic.Bool +} + +func (r *regionReq) finish() bool { + if !r.release() { + return false + } + + cost := time.Since(r.createTime).Seconds() + if cost > 0 && cost < abnormalRequestDurationInSec { + log.Debug("cdc resolve region request", + zap.Uint64("subID", uint64(r.regionInfo.subscribedSpan.subID)), + zap.Uint64("regionID", r.regionInfo.verID.GetID()), + zap.Float64("cost", cost), + zap.Int("inflightCount", r.controller.stats().inflight)) + metrics.RegionRequestFinishScanDuration.Observe(cost) + return true + } + log.Info("region request duration abnormal, skip metric", + zap.Float64("cost", cost), + zap.Uint64("regionID", r.regionInfo.verID.GetID())) + return true +} + +func (r *regionReq) abort() bool { + return r.release() +} + +func (r *regionReq) isActive() bool { + return !r.released.Load() +} + +func (r *regionReq) release() bool { + if !r.released.CompareAndSwap(false, true) { + return false + } + r.controller.release(r.scanBytes) + return true +} + +// regionAdmissionController owns one request worker's pending queue and +// initial-scan window. +type regionAdmissionController struct { + // mu guards the window state, inflight count, pending queue and closed flag. + mu sync.Mutex + + // currentWindow limits ordinary region scans. + currentWindow int + // maxWindow is the hard limit for high-priority region scans. + maxWindow int + // inflight is the number of admitted regions that have not finished their + // initial scan. It is guarded by mu. + inflight int + // pending keeps requests that have not entered the initial-scan window. + // It is guarded by mu. + pending *heap.Heap[*regionPriorityTask] + // memoryQuota gates initial scans using the log puller's global memory + // pressure. pdClock provides the current TS used for scan estimation. + memoryQuota *memoryQuotaController + pdClock pdutil.Clock + // notify wakes workers when a request is submitted or an admission slot is + // released. The one-element buffer prevents a wakeup from being lost between + // checking the admission condition and waiting on this channel. Notifications + // are only signals to recheck state; they do not correspond one-to-one with + // pending requests or available slots. + notify chan struct{} + // closed prevents new submissions and makes waiting workers exit. + closed bool +} + +type regionAdmissionStats struct { + pending int + inflight int +} + +func newRegionAdmissionController( + currentWindow int, + maxWindowMultiplier int, + memoryQuota *memoryQuotaController, + pdClock pdutil.Clock, +) *regionAdmissionController { + if currentWindow <= 0 { + currentWindow = 1 + } + if maxWindowMultiplier <= 0 { + maxWindowMultiplier = 1 + } + maxWindow := math.MaxInt + if currentWindow <= math.MaxInt/maxWindowMultiplier { + maxWindow = currentWindow * maxWindowMultiplier + } + return ®ionAdmissionController{ + currentWindow: currentWindow, + maxWindow: maxWindow, + pending: heap.NewHeap[*regionPriorityTask](), + memoryQuota: memoryQuota, + pdClock: pdClock, + notify: make(chan struct{}, 1), + } +} + +func (c *regionAdmissionController) submit(task *regionPriorityTask) bool { + c.mu.Lock() + if c.closed { + c.mu.Unlock() + return false + } + c.pending.AddOrUpdate(task) + c.notifyOneLocked() + c.mu.Unlock() + return true +} + +// pop waits for an eligible request. If interrupt is signaled first, it +// returns nil without an error so the worker can handle the interrupt source. +func (c *regionAdmissionController) pop( + ctx context.Context, + interrupt <-chan struct{}, +) (*regionReq, error) { + for { + c.mu.Lock() + if c.closed { + c.mu.Unlock() + return nil, context.Canceled + } + request, scanBytes, memoryReady := c.popEligibleLocked() + if request != nil { + c.inflight++ + c.mu.Unlock() + return ®ionReq{ + regionInfo: request.regionInfo, + createTime: time.Now(), + controller: c, + scanBytes: scanBytes, + }, nil + } + c.mu.Unlock() + + waitingForMemory := memoryReady != nil + if waitingForMemory { + c.memoryQuota.scanWaiters.Add(1) + } + waitStart := time.Time{} + if waitingForMemory { + waitStart = time.Now() + } + select { + case <-c.notify: + case <-memoryReady: + case <-interrupt: + if waitingForMemory { + c.memoryQuota.scanWaiters.Add(-1) + metrics.LogPullerMemoryQuotaScanWaitDuration.Observe(time.Since(waitStart).Seconds()) + } + return nil, nil + case <-ctx.Done(): + if waitingForMemory { + c.memoryQuota.scanWaiters.Add(-1) + metrics.LogPullerMemoryQuotaScanWaitDuration.Observe(time.Since(waitStart).Seconds()) + } + return nil, ctx.Err() + } + if waitingForMemory { + c.memoryQuota.scanWaiters.Add(-1) + metrics.LogPullerMemoryQuotaScanWaitDuration.Observe(time.Since(waitStart).Seconds()) + } + } +} + +func (c *regionAdmissionController) popEligibleLocked() ( + *regionPriorityTask, + uint64, + <-chan struct{}, +) { + request, ok := c.pending.PeekTop() + if !ok { + return nil, 0, nil + } + if c.inflight >= c.windowFor(request) { + return nil, 0, nil + } + + scanBytes, memoryReady, admitted := c.memoryQuota.AcquireScan( + request.regionInfo, c.pdClock.CurrentTS()) + if !admitted { + return nil, 0, memoryReady + } + request, _ = c.pending.PopTop() + return request, scanBytes, nil +} + +func (c *regionAdmissionController) windowFor(request *regionPriorityTask) int { + if request.canUseMaxWindow() { + return c.maxWindow + } + return c.currentWindow +} + +func (c *regionAdmissionController) release(scanBytes uint64) { + c.memoryQuota.ReleaseScan(scanBytes) + c.mu.Lock() + if c.inflight > 0 { + c.inflight-- + c.notifyOneLocked() + } + c.mu.Unlock() +} + +func (c *regionAdmissionController) close() { + c.mu.Lock() + if !c.closed { + c.closed = true + close(c.notify) + } + c.mu.Unlock() +} + +func (c *regionAdmissionController) stats() regionAdmissionStats { + c.mu.Lock() + defer c.mu.Unlock() + return regionAdmissionStats{ + pending: c.pending.Len(), + inflight: c.inflight, + } +} + +func (c *regionAdmissionController) drain() []*regionPriorityTask { + c.mu.Lock() + defer c.mu.Unlock() + + requests := make([]*regionPriorityTask, 0, c.pending.Len()) + for { + request, ok := c.pending.PopTop() + if !ok { + return requests + } + requests = append(requests, request) + } +} + +func (c *regionAdmissionController) notifyOneLocked() { + if c.closed { + return + } + select { + case c.notify <- struct{}{}: + default: + } +} diff --git a/logservice/logpuller/region_admission_controller_test.go b/logservice/logpuller/region_admission_controller_test.go new file mode 100644 index 0000000000..2633e99681 --- /dev/null +++ b/logservice/logpuller/region_admission_controller_test.go @@ -0,0 +1,224 @@ +// Copyright 2026 PingCAP, Inc. +// +// 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 logpuller + +import ( + "context" + "sync" + "testing" + "time" + + "github.com/pingcap/kvproto/pkg/cdcpb" + "github.com/pingcap/ticdc/heartbeatpb" + "github.com/pingcap/ticdc/logservice/logpuller/regionlock" + "github.com/pingcap/ticdc/pkg/pdutil" + "github.com/stretchr/testify/require" + "github.com/tikv/client-go/v2/oracle" + "github.com/tikv/client-go/v2/tikv" +) + +func createTestRegionInfo(subID SubscriptionID, regionID uint64) regionInfo { + span := heartbeatpb.TableSpan{ + TableID: 1, + StartKey: []byte("start"), + EndKey: []byte("end"), + } + return newRegionInfo( + tikv.NewRegionVerID(regionID, 1, 1), + span, + nil, + &subscribedSpan{subID: subID, startTs: 100, span: span}, + false, + ) +} + +func prepareRegionForAdmission(region regionInfo, checkpointTs uint64) regionInfo { + region.lockedRangeState = ®ionlock.LockedRangeState{} + region.lockedRangeState.ResolvedTs.Store(checkpointTs) + return region +} + +func submitRegionForAdmission( + t *testing.T, + controller *regionAdmissionController, + region regionInfo, + currentTs uint64, +) { + t.Helper() + task := newRegionPriorityTask(region, region.verID.GetID()) + require.True(t, controller.submit(task)) +} + +func newTestRegionAdmissionController( + currentWindow int, + maxWindowMultiplier int, +) *regionAdmissionController { + clock := pdutil.NewClock4Test() + return newRegionAdmissionController( + currentWindow, + maxWindowMultiplier, + newMemoryQuotaController(1024*1024*1024, 8*1024*1024), + clock, + ) +} + +func TestRegionAdmissionControllerNormalWindow(t *testing.T) { + controller := newTestRegionAdmissionController(1, 2) + currentTs := oracle.GoTimeToTS(time.Now()) + checkpointTs := oracle.GoTimeToTS(time.Now().Add(-time.Hour)) + region1 := prepareRegionForAdmission(createTestRegionInfo(1, 1), checkpointTs) + region2 := prepareRegionForAdmission(createTestRegionInfo(1, 2), checkpointTs) + submitRegionForAdmission(t, controller, region1, currentTs) + submitRegionForAdmission(t, controller, region2, currentTs) + + req1, err := controller.pop(t.Context(), nil) + require.NoError(t, err) + require.Equal(t, 1, controller.stats().inflight) + interrupt := make(chan struct{}) + close(interrupt) + req2, err := controller.pop(t.Context(), interrupt) + require.Nil(t, req2) + require.NoError(t, err) + + require.True(t, req1.abort()) + req2, err = controller.pop(t.Context(), nil) + require.NoError(t, err) + require.Equal(t, uint64(2), req2.regionInfo.verID.GetID()) + require.True(t, req2.abort()) +} + +func TestRegionAdmissionControllerHighPriorityUsesMaxWindow(t *testing.T) { + controller := newTestRegionAdmissionController(1, 2) + currentTs := oracle.GoTimeToTS(time.Now()) + slowCheckpointTs := oracle.GoTimeToTS(time.Now().Add(-time.Hour)) + + submitRegionForAdmission(t, controller, + prepareRegionForAdmission(createTestRegionInfo(1, 1), slowCheckpointTs), + currentTs) + req1, err := controller.pop(t.Context(), nil) + require.NoError(t, err) + + submitRegionForAdmission(t, controller, + prepareRegionForAdmission(createTestRegionInfo(1, 2), slowCheckpointTs), + currentTs) + highPriorityRegion := prepareRegionForAdmission(createTestRegionInfo(1, 3), slowCheckpointTs) + highPriorityRegion.scanPriority = cdcpb.ScanPriority_SCAN_PRIORITY_HIGH + submitRegionForAdmission(t, controller, highPriorityRegion, currentTs) + + req2, err := controller.pop(t.Context(), nil) + require.NoError(t, err) + require.Equal(t, uint64(3), req2.regionInfo.verID.GetID()) + interrupt := make(chan struct{}) + close(interrupt) + req3, err := controller.pop(t.Context(), interrupt) + require.Nil(t, req3) + require.NoError(t, err) + require.Equal(t, 2, controller.stats().inflight) + + require.True(t, req1.abort()) + require.True(t, req2.abort()) + req3, err = controller.pop(t.Context(), nil) + require.NoError(t, err) + require.Equal(t, uint64(2), req3.regionInfo.verID.GetID()) + require.True(t, req3.abort()) +} + +func TestRegionAdmissionControllerPrioritizesHighPriorityRegion(t *testing.T) { + controller := newTestRegionAdmissionController(1, 2) + currentTs := oracle.GoTimeToTS(time.Now()) + slowCheckpointTs := oracle.GoTimeToTS(time.Now().Add(-time.Hour)) + + submitRegionForAdmission(t, controller, + prepareRegionForAdmission(createTestRegionInfo(1, 1), slowCheckpointTs), + currentTs) + req1, err := controller.pop(t.Context(), nil) + require.NoError(t, err) + + submitRegionForAdmission(t, controller, + prepareRegionForAdmission(createTestRegionInfo(1, 2), slowCheckpointTs), + currentTs) + highPriorityRegion := prepareRegionForAdmission(createTestRegionInfo(1, 3), slowCheckpointTs) + highPriorityRegion.scanPriority = cdcpb.ScanPriority_SCAN_PRIORITY_HIGH + submitRegionForAdmission(t, controller, highPriorityRegion, currentTs) + + req2, err := controller.pop(t.Context(), nil) + require.NoError(t, err) + require.Equal(t, uint64(3), req2.regionInfo.verID.GetID()) + + require.True(t, req1.abort()) + require.True(t, req2.abort()) + req3, err := controller.pop(t.Context(), nil) + require.NoError(t, err) + require.Equal(t, uint64(2), req3.regionInfo.verID.GetID()) + require.True(t, req3.abort()) +} + +func TestRegionAdmissionLeaseReleasedOnce(t *testing.T) { + controller := newTestRegionAdmissionController(1, 1) + currentTs := oracle.GoTimeToTS(time.Now()) + region := prepareRegionForAdmission(createTestRegionInfo(1, 1), currentTs) + submitRegionForAdmission(t, controller, region, currentTs) + req, err := controller.pop(t.Context(), nil) + require.NoError(t, err) + + start := make(chan struct{}) + results := make(chan bool, 2) + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + <-start + results <- req.finish() + }() + go func() { + defer wg.Done() + <-start + results <- req.abort() + }() + close(start) + wg.Wait() + close(results) + + successes := 0 + for result := range results { + if result { + successes++ + } + } + require.Equal(t, 1, successes) + require.Zero(t, controller.stats().inflight) +} + +func TestRegionAdmissionControllerClose(t *testing.T) { + controller := newTestRegionAdmissionController(1, 1) + controller.close() + region := prepareRegionForAdmission(createTestRegionInfo(1, 1), 1) + require.False(t, controller.submit(newRegionPriorityTask(region, 1))) + + _, err := controller.pop(context.Background(), nil) + require.ErrorIs(t, err, context.Canceled) +} + +func TestRegionAdmissionControllerDrainPending(t *testing.T) { + controller := newTestRegionAdmissionController(1, 1) + region1 := prepareRegionForAdmission(createTestRegionInfo(1, 1), 1) + region2 := prepareRegionForAdmission(createTestRegionInfo(1, 2), 1) + submitRegionForAdmission(t, controller, region1, 1) + submitRegionForAdmission(t, controller, region2, 1) + + pending := controller.drain() + require.Len(t, pending, 2) + require.Zero(t, controller.stats().pending) +} diff --git a/logservice/logpuller/region_event_handler.go b/logservice/logpuller/region_event_handler.go index 0973b1df15..49dab3e46c 100644 --- a/logservice/logpuller/region_event_handler.go +++ b/logservice/logpuller/region_event_handler.go @@ -29,8 +29,7 @@ import ( ) var ( - metricsResolvedTsCount = metrics.PullerEventCounter.WithLabelValues("resolved_ts") - metricsEventCount = metrics.PullerEventCounter.WithLabelValues("event") + metricsEventCount = metrics.PullerEventCounter.WithLabelValues("event") metricRegionEventHandleDurationEntries = metrics.SubscriptionClientRegionEventHandleDuration.WithLabelValues("entries") metricRegionEventHandleDurationResolved = metrics.SubscriptionClientRegionEventHandleDuration.WithLabelValues("resolved") @@ -56,6 +55,13 @@ type regionEvent struct { entries *cdcpb.Event_Entries_ resolvedTs uint64 + // memoryBytes is released when this event is dropped or when the derived KV + // events no longer need to be retained by the log puller. + memoryBytes uint64 +} + +func (event *regionEvent) needsMemoryAccounting() bool { + return event.entries != nil } func (event *regionEvent) getSize() int { @@ -122,7 +128,9 @@ func (h *regionEventHandler) Handle(span *subscribedSpan, events ...regionEvent) } newResolvedTs := uint64(0) + memoryBytes := uint64(0) for _, event := range events { + memoryBytes += event.memoryBytes if len(event.states) == 1 && event.states[0].isStale() { hasError = true h.handleRegionError(event.states[0]) @@ -148,9 +156,13 @@ func (h *regionEventHandler) Handle(span *subscribedSpan, events ...regionEvent) span.advanceResolvedTs(newResolvedTs) } } + releaseMemoryQuota := func() { + h.eventSink.memoryQuota.ReleaseEvent(memoryBytes) + } if len(span.kvEventsCache) > 0 { metricsEventCount.Add(float64(len(span.kvEventsCache))) await := span.consumeKVEvents(span.kvEventsCache, func() { + defer releaseMemoryQuota() start := time.Now() span.clearKVEventsCache() metricConsumeKVEventsCallbackDurationClearCache.Observe(time.Since(start).Seconds()) @@ -167,10 +179,12 @@ func (h *regionEventHandler) Handle(span *subscribedSpan, events ...regionEvent) if !await { span.clearKVEventsCache() tryAdvanceResolvedTs() + releaseMemoryQuota() } return await } else { tryAdvanceResolvedTs() + releaseMemoryQuota() } return false } @@ -227,6 +241,7 @@ func (h *regionEventHandler) GetType(event regionEvent) dynstream.EventType { } func (h *regionEventHandler) OnDrop(event regionEvent) interface{} { + h.eventSink.memoryQuota.ReleaseEvent(event.memoryBytes) // TODO: Distinguish between drop events caused by "path not found" errors and memory control. state := event.mustFirstState() fields := []zap.Field{ @@ -255,7 +270,7 @@ func (h *regionEventHandler) handleRegionError(state *regionFeedState) { zap.Error(err)) } if stepsToRemoved { - worker.takeRegionState(SubscriptionID(state.requestID), state.getRegionID()) + worker.tracker.RemoveIf(SubscriptionID(state.requestID), state.getRegionID(), state) h.failureHandler.Report(newRegionErrorInfo(state.getRegionInfo(), err)) } } diff --git a/logservice/logpuller/region_event_handler_test.go b/logservice/logpuller/region_event_handler_test.go index 3a5b250cef..bbdd5cbee4 100644 --- a/logservice/logpuller/region_event_handler_test.go +++ b/logservice/logpuller/region_event_handler_test.go @@ -49,7 +49,10 @@ import ( func TestHandleEventEntryEventOutOfOrder(t *testing.T) { // initialize option := dynstream.NewOption() - ds := dynstream.NewParallelDynamicStream("test", ®ionEventHandler{}, option) + handler := ®ionEventHandler{eventSink: ®ionEventSink{ + memoryQuota: newMemoryQuotaController(1024*1024*1024, 8*1024*1024), + }} + ds := dynstream.NewParallelDynamicStream("test", handler, option) ds.Start() span := heartbeatpb.TableSpan{ @@ -72,7 +75,6 @@ func TestHandleEventEntryEventOutOfOrder(t *testing.T) { subID: subID, span: span, startTs: 1000, // not used - rangeLock: regionlock.NewRangeLock(uint64(subID), span.StartKey, span.EndKey, 1000), consumeKVEvents: consumeKVEvents, advanceResolvedTs: advanceResolvedTs, advanceInterval: 0, @@ -80,21 +82,17 @@ func TestHandleEventEntryEventOutOfOrder(t *testing.T) { ds.AddPath(subID, subSpan, dynstream.AreaSettings{}) worker := ®ionRequestWorker{ - requestCache: &requestCache{}, + tracker: newRegionTracker(), } region := newRegionInfo( - tikv.NewRegionVerID(1, 1, 1), + tikv.RegionVerID{}, span, &tikv.RPCContext{}, subSpan, false, ) - lockResult := subSpan.rangeLock.LockRange( - context.Background(), span.StartKey, span.EndKey, 1, 1) - require.Equal(t, regionlock.LockRangeStatusSuccess, lockResult.Status) - region.lockedRangeState = lockResult.LockedRangeState - state := newRegionFeedState(region, 1, worker) - state.start() + region.lockedRangeState = ®ionlock.LockedRangeState{} + state := newRegionFeedState(region, 1, worker, nil, nil) // Receive prewrite2 with empty value. { @@ -211,9 +209,10 @@ func TestHandleEventEntryEventOutOfOrder(t *testing.T) { func TestHandleResolvedTs(t *testing.T) { // initialize option := dynstream.NewOption() - pdClock := pdutil.NewClock4Test() - pdClock.(*pdutil.Clock4Test).SetTS(10) - ds := dynstream.NewParallelDynamicStream("test", ®ionEventHandler{}, option) + handler := ®ionEventHandler{eventSink: ®ionEventSink{ + memoryQuota: newMemoryQuotaController(1024*1024*1024, 8*1024*1024), + }} + ds := dynstream.NewParallelDynamicStream("test", handler, option) ds.Start() consumeKVEvents := func(events []common.RawKVEntry, _ func()) bool { return false } // not used @@ -224,39 +223,33 @@ func TestHandleResolvedTs(t *testing.T) { subID1 := SubscriptionID(1) worker := ®ionRequestWorker{ - requestCache: &requestCache{}, + tracker: newRegionTracker(), } - state1 := newRegionFeedState(regionInfo{verID: tikv.NewRegionVerID(1, 1, 1)}, uint64(subID1), worker) - state1.start() - var subSpan1 *subscribedSpan + state1 := newRegionFeedState(regionInfo{verID: tikv.NewRegionVerID(1, 1, 1)}, uint64(subID1), worker, nil, nil) { span := heartbeatpb.TableSpan{ TableID: 100, StartKey: common.ToComparableKey([]byte{}), // TODO: remove spanz dependency EndKey: common.ToComparableKey(common.UpperBoundKey), } - subSpan1 = &subscribedSpan{ + subSpan := &subscribedSpan{ subID: subID1, span: heartbeatpb.TableSpan{}, rangeLock: regionlock.NewRangeLock(uint64(subID1), span.StartKey, span.EndKey, 1), consumeKVEvents: consumeKVEvents, advanceResolvedTs: advanceResolvedTs, advanceInterval: 0, - priorityPolicy: newScanPriorityPolicy(pdClock, 30*time.Minute), + priorityPolicy: newScanPriorityPolicy(pdutil.NewClock4Test(), 30*time.Minute), } - ds.AddPath(subID1, subSpan1, dynstream.AreaSettings{}) - state1.region.subscribedSpan = subSpan1 - lockResult := subSpan1.rangeLock.LockRange( - context.Background(), span.StartKey, span.EndKey, 1, 1) - require.Equal(t, regionlock.LockRangeStatusSuccess, lockResult.Status) - state1.region.lockedRangeState = lockResult.LockedRangeState + ds.AddPath(subID1, subSpan, dynstream.AreaSettings{}) + state1.region.subscribedSpan = subSpan + state1.region.lockedRangeState = ®ionlock.LockedRangeState{} state1.setInitialized() state1.updateResolvedTs(9) } subID2 := SubscriptionID(2) - state2 := newRegionFeedState(regionInfo{verID: tikv.NewRegionVerID(2, 2, 2)}, uint64(subID2), worker) - state2.start() + state2 := newRegionFeedState(regionInfo{verID: tikv.NewRegionVerID(2, 2, 2)}, uint64(subID2), worker, nil, nil) { span := heartbeatpb.TableSpan{ TableID: 100, @@ -270,21 +263,17 @@ func TestHandleResolvedTs(t *testing.T) { consumeKVEvents: consumeKVEvents, advanceResolvedTs: advanceResolvedTs, advanceInterval: 0, - priorityPolicy: newScanPriorityPolicy(pdClock, 30*time.Minute), + priorityPolicy: newScanPriorityPolicy(pdutil.NewClock4Test(), 30*time.Minute), } ds.AddPath(subID2, subSpan, dynstream.AreaSettings{}) state2.region.subscribedSpan = subSpan - lockResult := subSpan.rangeLock.LockRange( - context.Background(), span.StartKey, span.EndKey, 2, 2) - require.Equal(t, regionlock.LockRangeStatusSuccess, lockResult.Status) - state2.region.lockedRangeState = lockResult.LockedRangeState + state2.region.lockedRangeState = ®ionlock.LockedRangeState{} state2.setInitialized() state2.updateResolvedTs(11) } subID3 := SubscriptionID(3) - state3 := newRegionFeedState(regionInfo{verID: tikv.NewRegionVerID(3, 3, 3)}, uint64(subID3), worker) - state3.start() + state3 := newRegionFeedState(regionInfo{verID: tikv.NewRegionVerID(3, 3, 3)}, uint64(subID3), worker, nil, nil) { span := heartbeatpb.TableSpan{ TableID: 100, @@ -298,14 +287,11 @@ func TestHandleResolvedTs(t *testing.T) { consumeKVEvents: consumeKVEvents, advanceResolvedTs: advanceResolvedTs, advanceInterval: 0, - priorityPolicy: newScanPriorityPolicy(pdClock, 30*time.Minute), + priorityPolicy: newScanPriorityPolicy(pdutil.NewClock4Test(), 30*time.Minute), } ds.AddPath(subID3, subSpan, dynstream.AreaSettings{}) state3.region.subscribedSpan = subSpan - lockResult := subSpan.rangeLock.LockRange( - context.Background(), span.StartKey, span.EndKey, 3, 3) - require.Equal(t, regionlock.LockRangeStatusSuccess, lockResult.Status) - state3.region.lockedRangeState = lockResult.LockedRangeState + state3.region.lockedRangeState = ®ionlock.LockedRangeState{} state3.updateResolvedTs(8) } @@ -350,7 +336,6 @@ func TestHandleResolvedTs(t *testing.T) { require.Equal(t, uint64(10), state1.getLastResolvedTs()) require.Equal(t, uint64(11), state2.getLastResolvedTs()) require.Equal(t, uint64(8), state3.getLastResolvedTs()) - require.True(t, subSpan1.priorityPolicy.everCaughtUp.Load()) } func TestHandleResolvedTsThrottled(t *testing.T) { @@ -383,6 +368,7 @@ func TestHandleResolvedTsThrottled(t *testing.T) { priorityPolicy: newScanPriorityPolicy(pdutil.NewClock4Test(), 30*time.Minute), } span.lastAdvanceTime.Store(0) + worker := ®ionRequestWorker{tracker: newRegionTracker()} state := newRegionFeedState( regionInfo{ verID: tikv.NewRegionVerID(1, 1, 1), @@ -390,70 +376,129 @@ func TestHandleResolvedTsThrottled(t *testing.T) { lockedRangeState: res1.LockedRangeState, }, 1, + worker, + nil, nil, ) - state.start() require.Equal(t, uint64(200), handleResolvedTs(span, state, 300)) } -func TestSpanInitializedAfterAllRangesInitialized(t *testing.T) { - ctx := context.Background() - rangeLock := regionlock.NewRangeLock(1, []byte("a"), []byte("z"), 100) - firstLock := rangeLock.LockRange(ctx, []byte("a"), []byte("m"), 1, 1) - require.Equal(t, regionlock.LockRangeStatusSuccess, firstLock.Status) - secondLock := rangeLock.LockRange(ctx, []byte("m"), []byte("z"), 2, 1) - require.Equal(t, regionlock.LockRangeStatusSuccess, secondLock.Status) +func TestHandleEntriesReleasesMemoryAfterDownstreamCallback(t *testing.T) { + quota := newMemoryQuotaController(1024, 8) + span := newTestQuotaSpan(1) + callbackCh := make(chan func(), 1) + span.consumeKVEvents = func(_ []common.RawKVEntry, callback func()) bool { + callbackCh <- callback + return true + } + span.advanceResolvedTs = func(uint64) {} + + lockedState := ®ionlock.LockedRangeState{} + lockedState.ResolvedTs.Store(100) + state := ®ionFeedState{ + region: regionInfo{ + verID: tikv.NewRegionVerID(1, 1, 1), + rpcCtx: &tikv.RPCContext{}, + subscribedSpan: span, + lockedRangeState: lockedState, + }, + } + require.True(t, quota.AcquireEvent(context.Background(), span, 10)) + handler := ®ionEventHandler{eventSink: ®ionEventSink{ + ds: newMockRegionEventSinkStream(), + memoryQuota: quota, + }} + + await := handler.Handle(span, regionEvent{ + states: []*regionFeedState{state}, + memoryBytes: 10, + entries: &cdcpb.Event_Entries_{Entries: &cdcpb.Event_Entries{ + Entries: []*cdcpb.Event_Row{{ + Type: cdcpb.Event_COMMITTED, + OpType: cdcpb.Event_Row_PUT, + CommitTs: 101, + }}, + }}, + }) + require.True(t, await) + quotaState := getMemoryQuotaTestState(quota) + require.Equal(t, uint64(10), quotaState.used) + + callback := <-callbackCh + callback() + quotaState = getMemoryQuotaTestState(quota) + require.Zero(t, quotaState.used) +} +func TestRegionEventHandlerInitializedResetsRecoveryState(t *testing.T) { span := &subscribedSpan{ - subID: SubscriptionID(1), - startTs: 100, - span: heartbeatpb.TableSpan{StartKey: []byte("a"), EndKey: []byte("z")}, - rangeLock: rangeLock, - priorityPolicy: newScanPriorityPolicy(pdutil.NewClock4Test(), 30*time.Minute), + subID: 1, + span: heartbeatpb.TableSpan{TableID: 1}, + advanceResolvedTs: func(uint64) {}, } - span.resolvedTs.Store(span.startTs) - worker := ®ionRequestWorker{requestCache: newRequestCache(2)} - newState := func( - regionID uint64, regionSpan heartbeatpb.TableSpan, - lockedRangeState *regionlock.LockedRangeState, - ) *regionFeedState { - state := newRegionFeedState( - regionInfo{ - verID: tikv.NewRegionVerID(regionID, 1, 1), - span: regionSpan, - rpcCtx: &tikv.RPCContext{}, - subscribedSpan: span, - lockedRangeState: lockedRangeState, - }, - uint64(span.subID), - worker, - ) - state.start() - return state + failureHandler := newRegionFailureHandler(nil, func(*subscribedSpan) {}, func(context.Context, regionInfo) {}, func(context.Context, rangeTask) {}) + key := newRegionRecoveryKey(span.subID, span.span) + failureHandler.recoveries[key] = ®ionRecoveryState{} + + region := newRegionInfo(tikv.NewRegionVerID(1, 1, 1), span.span, nil, span, false) + region.lockedRangeState = ®ionlock.LockedRangeState{} + region.rpcCtx = &tikv.RPCContext{Addr: "store-1"} + state := newRegionFeedState(region, uint64(span.subID), ®ionRequestWorker{tracker: newRegionTracker()}, nil, func(state *regionFeedState) { + failureHandler.resetRegionRecovery(state.region) + }) + + handler := ®ionEventHandler{ + eventSink: ®ionEventSink{memoryQuota: newMemoryQuotaController(0, 0)}, + failureHandler: failureHandler, } - firstState := newState(1, - heartbeatpb.TableSpan{StartKey: []byte("a"), EndKey: []byte("m")}, - firstLock.LockedRangeState) - secondState := newState(2, - heartbeatpb.TableSpan{StartKey: []byte("m"), EndKey: []byte("z")}, - secondLock.LockedRangeState) - - handler := ®ionEventHandler{} - initializedEvent := func(state *regionFeedState) regionEvent { - return regionEvent{ - states: []*regionFeedState{state}, - entries: &cdcpb.Event_Entries_{Entries: &cdcpb.Event_Entries{ + handler.Handle(span, regionEvent{ + states: []*regionFeedState{state}, + entries: &cdcpb.Event_Entries_{ + Entries: &cdcpb.Event_Entries{ Entries: []*cdcpb.Event_Row{{Type: cdcpb.Event_INITIALIZED}}, - }}, - } + }, + }, + }) + + failureHandler.recoveryMu.Lock() + _, ok := failureHandler.recoveries[key] + failureHandler.recoveryMu.Unlock() + require.False(t, ok) +} + +func TestSpanInitializedAfterFullRangeCoverage(t *testing.T) { + const startTs = 100 + span := &subscribedSpan{ + subID: 1, + startTs: startTs, + span: heartbeatpb.TableSpan{ + StartKey: []byte("a"), + EndKey: []byte("z"), + }, } + firstState := newRegionFeedState(regionInfo{ + verID: tikv.NewRegionVerID(1, 1, 1), + span: heartbeatpb.TableSpan{ + StartKey: []byte("a"), + EndKey: []byte("m"), + }, + subscribedSpan: span, + lockedRangeState: ®ionlock.LockedRangeState{}, + }, uint64(span.subID), ®ionRequestWorker{}, nil, nil) + secondState := newRegionFeedState(regionInfo{ + verID: tikv.NewRegionVerID(2, 1, 1), + span: heartbeatpb.TableSpan{ + StartKey: []byte("m"), + EndKey: []byte("z"), + }, + subscribedSpan: span, + lockedRangeState: ®ionlock.LockedRangeState{}, + }, uint64(span.subID), ®ionRequestWorker{}, nil, nil) - require.False(t, handler.Handle(span, initializedEvent(firstState))) + span.markRegionInitialized(firstState) require.False(t, span.initialized.Load()) - require.Equal(t, uint64(0), handleResolvedTs(span, firstState, span.startTs)) - require.False(t, handler.Handle(span, initializedEvent(secondState))) + span.markRegionInitialized(secondState) require.True(t, span.initialized.Load()) - require.Equal(t, span.startTs, handleResolvedTs(span, secondState, span.startTs)) } diff --git a/logservice/logpuller/region_event_sink.go b/logservice/logpuller/region_event_sink.go index a89a0555bc..64c837f320 100644 --- a/logservice/logpuller/region_event_sink.go +++ b/logservice/logpuller/region_event_sink.go @@ -15,32 +15,26 @@ package logpuller import ( "context" - "sync" - "sync/atomic" "github.com/pingcap/log" - "github.com/pingcap/ticdc/pkg/metrics" "github.com/pingcap/ticdc/utils/dynstream" "go.uber.org/zap" ) -// regionEventSink delivers region events to dynstream and owns push-side flow control. +// regionEventSink delivers region events to dynstream and accounts their memory. type regionEventSink struct { - // mu/cond coordinate the paused push path with pause/resume and shutdown signals. - mu sync.Mutex - cond *sync.Cond - // paused tracks whether region event pushing is temporarily held back by feedback. - paused atomic.Bool - // stopped marks the sink as shutting down so blocked pushers can exit instead of waiting for resume. - stopped atomic.Bool + ctx context.Context + ds dynstream.DynamicStream[int, SubscriptionID, regionEvent, *subscribedSpan, *regionEventHandler] - // ds owns the dynstream used to deliver region events and receive flow-control feedback. - ds dynstream.DynamicStream[int, SubscriptionID, regionEvent, *subscribedSpan, *regionEventHandler] + memoryQuota *memoryQuotaController } -func newRegionEventSink(failureHandler *regionFailureHandler) *regionEventSink { - sink := ®ionEventSink{} - sink.cond = sync.NewCond(&sink.mu) +func newRegionEventSink( + ctx context.Context, + failureHandler *regionFailureHandler, + memoryQuota *memoryQuotaController, +) *regionEventSink { + sink := ®ionEventSink{ctx: ctx, memoryQuota: memoryQuota} option := dynstream.NewOption() // Note: it is max batch size of the kv sent from tikv(not committed rows) @@ -48,7 +42,6 @@ func newRegionEventSink(failureHandler *regionFailureHandler) *regionEventSink { // TODO: Set `UseBuffer` to true until we refactor the `regionEventHandler.Handle` method so that it doesn't call any method of the dynamic stream. Currently, if `UseBuffer` is set to false, there will be a deadlock: // ds.handleLoop fetch events from `ch` -> regionEventHandler.Handle -> ds.RemovePath -> send event to `ch` option.UseBuffer = true - option.EnableMemoryControl = true ds := dynstream.NewParallelDynamicStream( "log-puller", ®ionEventHandler{eventSink: sink, failureHandler: failureHandler}, @@ -60,8 +53,7 @@ func newRegionEventSink(failureHandler *regionFailureHandler) *regionEventSink { } func (s *regionEventSink) AddPath(rt *subscribedSpan) { - areaSetting := dynstream.NewAreaSettingsWithMaxPendingSize(1*1024*1024*1024, dynstream.MemoryControlForPuller, "logPuller") // 1GB - if err := s.ds.AddPath(rt.subID, rt, areaSetting); err != nil { + if err := s.ds.AddPath(rt.subID, rt); err != nil { log.Warn("subscription client add path failed", zap.Uint64("subscriptionID", uint64(rt.subID)), zap.Error(err)) @@ -77,108 +69,25 @@ func (s *regionEventSink) Wake(subID SubscriptionID) { } func (s *regionEventSink) Push(subID SubscriptionID, event regionEvent) { - if s.stopped.Load() { - return - } - // fast path - if !s.paused.Load() { - s.ds.Push(subID, event) - return - } - - // slow path: wait until paused is false - s.mu.Lock() - for s.paused.Load() && !s.stopped.Load() { - s.cond.Wait() - } - stopped := s.stopped.Load() - s.mu.Unlock() - - if stopped { - return - } - s.ds.Push(subID, event) -} - -func (s *regionEventSink) Run(ctx context.Context) error { - for { - select { - case <-ctx.Done(): - s.stop() - return nil - case feedback := <-s.ds.Feedback(): - switch feedback.FeedbackType { - case dynstream.PauseArea: - s.pause() - log.Info("subscription client pause push region event") - case dynstream.ResumeArea: - s.resume() - log.Info("subscription client resume push region event") - case dynstream.ReleasePath, dynstream.ResumePath: - // Ignore it, because it is no need to pause and resume a path in puller. - } + if event.needsMemoryAccounting() { + span := event.mustFirstState().region.subscribedSpan + event.memoryBytes = uint64(event.getSize()) + // AcquireEvent only returns false after shutdown or when the + // subscription has already been stopped. + if !s.memoryQuota.AcquireEvent(s.ctx, span, event.memoryBytes) { + return } } + s.ds.Push(subID, event) } func (s *regionEventSink) UpdateMetrics() { dsMetrics := s.ds.GetMetrics() metricSubscriptionClientDSChannelSize.Set(float64(dsMetrics.EventChanSize)) metricSubscriptionClientDSPendingQueueLen.Set(float64(dsMetrics.PendingQueueLen)) - - if len(dsMetrics.MemoryControl.AreaMemoryMetrics) == 0 { - return - } - if len(dsMetrics.MemoryControl.AreaMemoryMetrics) != 1 { - log.Warn("subscription client should have exactly one area") - return - } - - areaMetric := dsMetrics.MemoryControl.AreaMemoryMetrics[0] - metrics.DynamicStreamMemoryUsage.WithLabelValues( - "log-puller", - "max", - "default", - "default", - ).Set(float64(areaMetric.MaxMemory())) - metrics.DynamicStreamMemoryUsage.WithLabelValues( - "log-puller", - "used", - "default", - "default", - ).Set(float64(areaMetric.MemoryUsage())) } func (s *regionEventSink) Close() { - s.stop() + s.memoryQuota.WakeAll() s.ds.Close() } - -func (s *regionEventSink) pause() { - s.mu.Lock() - defer s.mu.Unlock() - if s.stopped.Load() || s.paused.Load() { - return - } - s.paused.Store(true) -} - -func (s *regionEventSink) resume() { - s.mu.Lock() - defer s.mu.Unlock() - if !s.paused.Load() { - return - } - s.paused.Store(false) - s.cond.Broadcast() -} - -func (s *regionEventSink) stop() { - if !s.stopped.CompareAndSwap(false, true) { - return - } - s.mu.Lock() - s.paused.Store(false) - s.cond.Broadcast() - s.mu.Unlock() -} diff --git a/logservice/logpuller/region_event_sink_test.go b/logservice/logpuller/region_event_sink_test.go index 0698d3fb4a..291d4d867c 100644 --- a/logservice/logpuller/region_event_sink_test.go +++ b/logservice/logpuller/region_event_sink_test.go @@ -15,28 +15,22 @@ package logpuller import ( "context" - "sync" - "sync/atomic" "testing" - "time" - "github.com/pingcap/ticdc/pkg/metrics" + "github.com/pingcap/kvproto/pkg/cdcpb" "github.com/pingcap/ticdc/utils/dynstream" "github.com/prometheus/client_golang/prometheus/testutil" "github.com/stretchr/testify/require" ) type mockRegionEventSinkStream struct { - feedbackCh chan dynstream.Feedback[int, SubscriptionID, *subscribedSpan] - pushCount atomic.Int32 - pushCh chan struct{} - metrics dynstream.Metrics[int, SubscriptionID] + eventCh chan regionEvent + metrics dynstream.Metrics[int, SubscriptionID] } func newMockRegionEventSinkStream() *mockRegionEventSinkStream { return &mockRegionEventSinkStream{ - feedbackCh: make(chan dynstream.Feedback[int, SubscriptionID, *subscribedSpan], 2), - pushCh: make(chan struct{}, 1), + eventCh: make(chan regionEvent, 1), } } @@ -44,15 +38,14 @@ func (s *mockRegionEventSinkStream) Start() {} func (s *mockRegionEventSinkStream) Close() {} -func (s *mockRegionEventSinkStream) Push(_ SubscriptionID, _ regionEvent) { - s.pushCount.Add(1) - s.pushCh <- struct{}{} +func (s *mockRegionEventSinkStream) Push(_ SubscriptionID, event regionEvent) { + s.eventCh <- event } func (s *mockRegionEventSinkStream) Wake(_ SubscriptionID) {} func (s *mockRegionEventSinkStream) Feedback() <-chan dynstream.Feedback[int, SubscriptionID, *subscribedSpan] { - return s.feedbackCh + return nil } func (s *mockRegionEventSinkStream) AddPath(_ SubscriptionID, _ *subscribedSpan, _ ...dynstream.AreaSettings) error { @@ -71,191 +64,46 @@ func (s *mockRegionEventSinkStream) GetMetrics() dynstream.Metrics[int, Subscrip return s.metrics } -func newTestRegionEventSink( - ds dynstream.DynamicStream[int, SubscriptionID, regionEvent, *subscribedSpan, *regionEventHandler], -) *regionEventSink { - sink := ®ionEventSink{ds: ds} - sink.cond = sync.NewCond(&sink.mu) - return sink -} - -func TestRegionEventSinkRunPausesAndResumesPush(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - +func TestRegionEventSinkUpdateMetrics(t *testing.T) { ds := newMockRegionEventSinkStream() - sink := newTestRegionEventSink(ds) - - runErrCh := make(chan error, 1) - go func() { - runErrCh <- sink.Run(ctx) - }() - - ds.feedbackCh <- dynstream.Feedback[int, SubscriptionID, *subscribedSpan]{ - FeedbackType: dynstream.PauseArea, - } - require.Eventually(t, sink.paused.Load, time.Second, 10*time.Millisecond) - - pushDone := make(chan struct{}) - go func() { - sink.Push(SubscriptionID(1), regionEvent{resolvedTs: 100}) - close(pushDone) - }() - - select { - case <-pushDone: - t.Fatal("Push should block while the sink is paused") - case <-time.After(100 * time.Millisecond): - } - require.Equal(t, int32(0), ds.pushCount.Load()) - - ds.feedbackCh <- dynstream.Feedback[int, SubscriptionID, *subscribedSpan]{ - FeedbackType: dynstream.ResumeArea, - } - require.Eventually(t, func() bool { return !sink.paused.Load() }, time.Second, 10*time.Millisecond) - - select { - case <-ds.pushCh: - case <-time.After(time.Second): - t.Fatal("Push should resume after ResumeArea feedback") + ds.metrics = dynstream.Metrics[int, SubscriptionID]{ + EventChanSize: 33, + PendingQueueLen: 44, } - select { - case <-pushDone: - case <-time.After(time.Second): - t.Fatal("Push should return after ResumeArea feedback") - } - - cancel() - select { - case err := <-runErrCh: - require.NoError(t, err) - case <-time.After(time.Second): - t.Fatal("Run should exit after context cancellation") - } -} - -func TestRegionEventSinkUpdateMetrics(t *testing.T) { - t.Run("empty area metrics returns after queue gauges", func(t *testing.T) { - ds := newMockRegionEventSinkStream() - ds.metrics = dynstream.Metrics[int, SubscriptionID]{ - EventChanSize: 11, - PendingQueueLen: 22, - } - - metrics.DynamicStreamMemoryUsage.WithLabelValues( - "log-puller", - "max", - "default", - "default", - ).Set(123) - metrics.DynamicStreamMemoryUsage.WithLabelValues( - "log-puller", - "used", - "default", - "default", - ).Set(456) - - sink := ®ionEventSink{ - ds: ds, - } - sink.UpdateMetrics() - - require.Equal(t, float64(11), testutil.ToFloat64(metricSubscriptionClientDSChannelSize)) - require.Equal(t, float64(22), testutil.ToFloat64(metricSubscriptionClientDSPendingQueueLen)) - require.Equal(t, float64(123), testutil.ToFloat64(metrics.DynamicStreamMemoryUsage.WithLabelValues( - "log-puller", - "max", - "default", - "default", - ))) - require.Equal(t, float64(456), testutil.ToFloat64(metrics.DynamicStreamMemoryUsage.WithLabelValues( - "log-puller", - "used", - "default", - "default", - ))) - }) - - t.Run("single area metrics updates memory gauges", func(t *testing.T) { - ds := newMockRegionEventSinkStream() - ds.metrics = dynstream.Metrics[int, SubscriptionID]{ - EventChanSize: 33, - PendingQueueLen: 44, - MemoryControl: dynstream.MemoryMetric[int, SubscriptionID]{ - AreaMemoryMetrics: []dynstream.AreaMemoryMetric[int, SubscriptionID]{ - { - UsedMemoryValue: 55, - MaxMemoryValue: 66, - PathMaxMemoryValue: 66, - }, - }, - }, - } + sink := ®ionEventSink{ds: ds} - sink := ®ionEventSink{ - ds: ds, - } - sink.UpdateMetrics() + sink.UpdateMetrics() - require.Equal(t, float64(33), testutil.ToFloat64(metricSubscriptionClientDSChannelSize)) - require.Equal(t, float64(44), testutil.ToFloat64(metricSubscriptionClientDSPendingQueueLen)) - require.Equal(t, float64(66), testutil.ToFloat64(metrics.DynamicStreamMemoryUsage.WithLabelValues( - "log-puller", - "max", - "default", - "default", - ))) - require.Equal(t, float64(55), testutil.ToFloat64(metrics.DynamicStreamMemoryUsage.WithLabelValues( - "log-puller", - "used", - "default", - "default", - ))) - }) + require.Equal(t, float64(33), testutil.ToFloat64(metricSubscriptionClientDSChannelSize)) + require.Equal(t, float64(44), testutil.ToFloat64(metricSubscriptionClientDSPendingQueueLen)) } -func TestRegionEventSinkRunCancelUnblocksPush(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - - ds := newMockRegionEventSinkStream() - sink := newTestRegionEventSink(ds) - - runErrCh := make(chan error, 1) - go func() { - runErrCh <- sink.Run(ctx) - }() - - ds.feedbackCh <- dynstream.Feedback[int, SubscriptionID, *subscribedSpan]{ - FeedbackType: dynstream.PauseArea, - } - require.Eventually(t, sink.paused.Load, time.Second, 10*time.Millisecond) - - pushDone := make(chan struct{}) - go func() { - sink.Push(SubscriptionID(1), regionEvent{resolvedTs: 100}) - close(pushDone) - }() - - select { - case <-pushDone: - t.Fatal("Push should block while the sink is paused") - case <-time.After(100 * time.Millisecond): +func TestRegionEventSinkTracksEntriesUntilDrop(t *testing.T) { + quota := newMemoryQuotaController(1024, 8) + span := newTestQuotaSpan(1) + state := ®ionFeedState{ + region: regionInfo{subscribedSpan: span}, + worker: ®ionRequestWorker{}, } - require.Equal(t, int32(0), ds.pushCount.Load()) - - cancel() - - select { - case <-pushDone: - case <-time.After(time.Second): - t.Fatal("Push should be unblocked by Run context cancellation") + ds := newMockRegionEventSinkStream() + sink := ®ionEventSink{ + ctx: context.Background(), + ds: ds, + memoryQuota: quota, } - require.Equal(t, int32(0), ds.pushCount.Load()) - select { - case err := <-runErrCh: - require.NoError(t, err) - case <-time.After(time.Second): - t.Fatal("Run should exit after context cancellation") - } + sink.Push(span.subID, regionEvent{ + states: []*regionFeedState{state}, + entries: &cdcpb.Event_Entries_{Entries: &cdcpb.Event_Entries{ + Entries: []*cdcpb.Event_Row{{Key: []byte("key"), Value: []byte("value")}}, + }}, + }) + pushed := <-ds.eventCh + require.NotZero(t, pushed.memoryBytes) + quotaState := getMemoryQuotaTestState(quota) + require.NotZero(t, quotaState.used) + + (®ionEventHandler{eventSink: sink}).OnDrop(pushed) + quotaState = getMemoryQuotaTestState(quota) + require.Zero(t, quotaState.used) } diff --git a/logservice/logpuller/region_failure_handler.go b/logservice/logpuller/region_failure_handler.go index 55eee0975e..64568df09c 100644 --- a/logservice/logpuller/region_failure_handler.go +++ b/logservice/logpuller/region_failure_handler.go @@ -15,10 +15,13 @@ package logpuller import ( "context" + "math/rand/v2" "sync" "time" + "github.com/pingcap/kvproto/pkg/cdcpb" "github.com/pingcap/log" + "github.com/pingcap/ticdc/heartbeatpb" "github.com/pingcap/ticdc/pkg/errors" "github.com/pingcap/ticdc/pkg/metrics" "github.com/tikv/client-go/v2/tikv" @@ -40,15 +43,148 @@ var ( // regionFailureHandler handles failed regions and owns retry and reschedule decisions. type regionFailureHandler struct { - cache *errCache - client *subscriptionClient + cache *errCache + regionCache *tikv.RegionCache + recoveryMu sync.Mutex + recoveries map[regionRecoveryKey]*regionRecoveryState + + onTableDrained func(*subscribedSpan) + scheduleRegionRequest func(context.Context, regionInfo) + scheduleRangeRequest func(context.Context, rangeTask) +} + +const ( + regionRecoveryBaseDelay = 50 * time.Millisecond + regionRecoveryMaxDelay = 2 * time.Second + regionRecoveryStateTTL = 5 * time.Minute +) + +// regionRecoveryKey keeps backoff state across region ID and epoch changes for +// the same logical range. +type regionRecoveryKey struct { + subscriptionID SubscriptionID + startKey string + endKey string +} + +type regionRecoveryState struct { + attempt uint32 + expiresAt time.Time +} + +func newRegionRecoveryKey( + subscriptionID SubscriptionID, + span heartbeatpb.TableSpan, +) regionRecoveryKey { + return regionRecoveryKey{ + subscriptionID: subscriptionID, + startKey: string(span.StartKey), + endKey: string(span.EndKey), + } } -func newRegionFailureHandler(client *subscriptionClient) *regionFailureHandler { +func regionRecoveryDelay(attempt uint32) time.Duration { + if attempt == 0 { + attempt = 1 + } + exponent := attempt - 1 + if exponent > 16 { + exponent = 16 + } + delay := regionRecoveryBaseDelay << exponent + if delay > regionRecoveryMaxDelay { + delay = regionRecoveryMaxDelay + } + half := delay / 2 + return half + time.Duration(rand.Int64N(int64(delay-half)+1)) +} + +func newRegionFailureHandler( + regionCache *tikv.RegionCache, + onTableDrained func(*subscribedSpan), + scheduleRegionRequest func(context.Context, regionInfo), + scheduleRangeRequest func(context.Context, rangeTask), +) *regionFailureHandler { return ®ionFailureHandler{ - cache: newErrCache(), - client: client, + cache: newErrCache(), + regionCache: regionCache, + recoveries: make(map[regionRecoveryKey]*regionRecoveryState), + onTableDrained: onTableDrained, + scheduleRegionRequest: scheduleRegionRequest, + scheduleRangeRequest: scheduleRangeRequest, + } +} + +func (r *regionFailureHandler) scheduleRecovery( + ctx context.Context, + subscribedSpan *subscribedSpan, + span heartbeatpb.TableSpan, + minDelay time.Duration, + retry func(), +) { + if subscribedSpan == nil || subscribedSpan.stopped.Load() { + return + } + key := newRegionRecoveryKey(subscribedSpan.subID, span) + + r.recoveryMu.Lock() + state := r.recoveries[key] + if state == nil { + state = ®ionRecoveryState{} + r.recoveries[key] = state + } + if state.attempt < 32 { + state.attempt++ + } + delay := regionRecoveryDelay(state.attempt) + if minDelay > delay { + delay = minDelay } + state.expiresAt = time.Now().Add(delay + regionRecoveryStateTTL) + r.recoveryMu.Unlock() + + time.AfterFunc(delay, func() { + r.recoveryMu.Lock() + if r.recoveries[key] != state { + r.recoveryMu.Unlock() + return + } + // Keep the attempt until the retry succeeds or the state expires. + state.expiresAt = time.Now().Add(regionRecoveryStateTTL) + r.recoveryMu.Unlock() + + if ctx.Err() != nil || subscribedSpan.stopped.Load() { + r.resetRecovery(key) + return + } + retry() + }) +} + +func (r *regionFailureHandler) expireRecoveries(now time.Time) { + r.recoveryMu.Lock() + defer r.recoveryMu.Unlock() + for key, state := range r.recoveries { + if !state.expiresAt.After(now) { + delete(r.recoveries, key) + } + } +} + +func (r *regionFailureHandler) resetRecovery(key regionRecoveryKey) { + r.recoveryMu.Lock() + defer r.recoveryMu.Unlock() + delete(r.recoveries, key) +} + +func (r *regionFailureHandler) resetRegionRecovery(region regionInfo) { + r.resetRecovery(newRegionRecoveryKey(region.subscribedSpan.subID, region.span)) +} + +func (r *regionFailureHandler) cancelRecoveries() { + r.recoveryMu.Lock() + defer r.recoveryMu.Unlock() + clear(r.recoveries) } // Report admits a region failure into the recovery pipeline. It releases the @@ -58,13 +194,17 @@ func (r *regionFailureHandler) Report(errInfo regionErrorInfo) { if errInfo.subscribedSpan.rangeLock.UnlockRange( errInfo.span.StartKey, errInfo.span.EndKey, errInfo.verID.GetID(), errInfo.verID.GetVer(), errInfo.resolvedTs()) { - r.client.onTableDrained(errInfo.subscribedSpan) + r.onTableDrained(errInfo.subscribedSpan) return } r.cache.add(errInfo) } func (r *regionFailureHandler) Run(ctx context.Context) error { + log.Info("region failure handler starts") + defer log.Info("region failure handler exits") + defer r.cancelRecoveries() + handleCachedErrors := func() error { for { batch := r.cache.popBatch(errCacheBatchSize) @@ -86,14 +226,17 @@ func (r *regionFailureHandler) Run(ctx context.Context) error { // r.cache.ready() should handle failures promptly in normal flow. The ticker is only a // fallback scan and is not expected to be needed in practice. - ticker := time.NewTicker(200 * time.Millisecond) - defer ticker.Stop() + fallbackTicker := time.NewTicker(200 * time.Millisecond) + defer fallbackTicker.Stop() + cleanupTicker := time.NewTicker(regionRecoveryStateTTL) + defer cleanupTicker.Stop() for { select { case <-ctx.Done(): - log.Info("subscription client handle errors and exit") return ctx.Err() - case <-ticker.C: + case now := <-cleanupTicker.C: + r.expireRecoveries(now) + case <-fallbackTicker.C: if err := handleCachedErrors(); err != nil { return err } @@ -107,7 +250,43 @@ func (r *regionFailureHandler) Run(ctx context.Context) error { func (r *regionFailureHandler) handleError(ctx context.Context, errInfo regionErrorInfo) error { err := errors.Cause(errInfo.err) - retryPriority := taskTypeFromScanPriority(errInfo.scanPriority) + retryRegion := func(minDelay time.Duration) { + r.scheduleRecovery( + ctx, + errInfo.subscribedSpan, + errInfo.span, + minDelay, + func() { + r.scheduleRegionRequest(ctx, errInfo.regionInfo) + }, + ) + } + retryRange := func() { + priority := normalizeScanPriority(errInfo.scanPriority) + if priority == cdcpb.ScanPriority_SCAN_PRIORITY_LOW { + priority = errInfo.subscribedSpan.priorityPolicy.resolve( + priority, + errInfo.resolvedTs(), + errInfo.subscribedSpan.priorityPolicy.pdClock.CurrentTime(), + ) + } + task := rangeTask{ + span: errInfo.span, + subscribedSpan: errInfo.subscribedSpan, + filterLoop: errInfo.filterLoop, + priority: priority, + } + r.scheduleRecovery( + ctx, + task.subscribedSpan, + task.span, + 0, + func() { + r.scheduleRangeRequest(ctx, task) + }, + ) + } + //nolint:errorlint // converting large type switch to errors.As is a significant refactor if _, requestCancelled := err.(*requestCancelledErr); !requestCancelled { log.Debug("cdc region error", @@ -122,28 +301,34 @@ func (r *regionFailureHandler) handleError(ctx context.Context, errInfo regionEr innerErr := eerr.err if notLeader := innerErr.GetNotLeader(); notLeader != nil { metricFeedNotLeaderCounter.Inc() - r.client.regionCache.UpdateLeader(errInfo.verID, notLeader.GetLeader(), errInfo.rpcCtx.AccessIdx) - r.client.scheduleRegionRequest(ctx, errInfo.regionInfo, retryPriority) + leader := notLeader.GetLeader() + if leader == nil || leader.GetId() == 0 || leader.GetStoreId() == 0 || errInfo.rpcCtx == nil { + r.regionCache.InvalidateCachedRegion(errInfo.verID) + retryRange() + return nil + } + r.regionCache.UpdateLeader(errInfo.verID, leader, errInfo.rpcCtx.AccessIdx) + retryRegion(0) return nil } if innerErr.GetEpochNotMatch() != nil { metricFeedEpochNotMatchCounter.Inc() - r.client.scheduleRangeRequest(ctx, errInfo.span, errInfo.subscribedSpan, errInfo.filterLoop, retryPriority) + retryRange() return nil } if innerErr.GetRegionNotFound() != nil { metricFeedRegionNotFoundCounter.Inc() - r.client.scheduleRangeRequest(ctx, errInfo.span, errInfo.subscribedSpan, errInfo.filterLoop, retryPriority) + retryRange() return nil } if innerErr.GetCongested() != nil { metricKvCongestedCounter.Inc() - r.client.scheduleRegionRequest(ctx, errInfo.regionInfo, retryPriority) + retryRegion(0) return nil } - if innerErr.GetServerIsBusy() != nil { + if busy := innerErr.GetServerIsBusy(); busy != nil { metricKvIsBusyCounter.Inc() - r.client.scheduleRegionRequest(ctx, errInfo.regionInfo, retryPriority) + retryRegion(time.Duration(busy.GetBackoffMs()) * time.Millisecond) return nil } if duplicated := innerErr.GetDuplicateRequest(); duplicated != nil { @@ -162,31 +347,34 @@ func (r *regionFailureHandler) handleError(ctx context.Context, errInfo regionEr zap.Uint64("subscriptionID", uint64(errInfo.subscribedSpan.subID)), zap.Stringer("error", innerErr)) metricFeedUnknownErrorCounter.Inc() - r.client.scheduleRegionRequest(ctx, errInfo.regionInfo, retryPriority) + retryRegion(0) return nil case *rpcCtxUnavailableErr: metricFeedRPCCtxUnavailable.Inc() - r.client.scheduleRangeRequest(ctx, errInfo.span, errInfo.subscribedSpan, errInfo.filterLoop, retryPriority) + retryRange() return nil case *getStoreErr: metricGetStoreErr.Inc() bo := tikv.NewBackoffer(ctx, tikvRequestMaxBackoff) // cannot get the store the region belongs to, so we need to reload the region. - r.client.regionCache.OnSendFail(bo, errInfo.rpcCtx, true, err) - r.client.scheduleRangeRequest(ctx, errInfo.span, errInfo.subscribedSpan, errInfo.filterLoop, retryPriority) + r.regionCache.OnSendFail(bo, errInfo.rpcCtx, true, err) + retryRange() return nil case *storeStreamErr: metricStoreSendRequestErr.Inc() bo := tikv.NewBackoffer(ctx, tikvRequestMaxBackoff) - r.client.regionCache.OnSendFail(bo, errInfo.rpcCtx, regionScheduleReload, err) - r.client.scheduleRegionRequest(ctx, errInfo.regionInfo, retryPriority) + r.regionCache.OnSendFail(bo, errInfo.rpcCtx, regionScheduleReload, err) + retryRegion(0) return nil case *requestCancelledErr: // the corresponding subscription has been unsubscribed, just ignore. + if errInfo.subscribedSpan != nil { + r.resetRegionRecovery(errInfo.regionInfo) + } return nil default: // TODO(qupeng): for some errors it's better to just deregister the region from TiKVs. - log.Warn("subscription client meets an internal error, fail the changefeed", + log.Warn("region failure cannot be recovered, fail the changefeed", zap.Uint64("subscriptionID", uint64(errInfo.subscribedSpan.subID)), zap.Error(err)) return err diff --git a/logservice/logpuller/region_failure_handler_test.go b/logservice/logpuller/region_failure_handler_test.go index 7f419920a7..0ce52a44a5 100644 --- a/logservice/logpuller/region_failure_handler_test.go +++ b/logservice/logpuller/region_failure_handler_test.go @@ -19,7 +19,11 @@ import ( "time" "github.com/pingcap/errors" + "github.com/pingcap/kvproto/pkg/cdcpb" + "github.com/pingcap/kvproto/pkg/errorpb" "github.com/pingcap/ticdc/heartbeatpb" + "github.com/pingcap/ticdc/pkg/pdutil" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/tikv/client-go/v2/tikv" ) @@ -115,7 +119,7 @@ func TestErrCachePopBatch(t *testing.T) { } func TestRegionFailureHandlerRunDrainsErrCacheWithoutDispatcher(t *testing.T) { - handler := newRegionFailureHandler(&subscriptionClient{}) + handler := newRegionFailureHandler(nil, func(*subscribedSpan) {}, func(context.Context, regionInfo) {}, func(context.Context, rangeTask) {}) for i := 0; i < errCacheBatchSize+5; i++ { handler.cache.add(newTestRegionErrorInfo(&requestCancelledErr{})) } @@ -142,3 +146,118 @@ func TestRegionFailureHandlerRunDrainsErrCacheWithoutDispatcher(t *testing.T) { t.Fatal("failure handler did not exit after context cancellation") } } + +func TestRegionFailureHandlerSchedulesNotLeaderRangeRetry(t *testing.T) { + pdClient := newFailureRecoveryTestPDClient(t) + defer pdClient.Close() + + regionCache := tikv.NewRegionCache(pdClient) + defer regionCache.Close() + + region := createFailureRecoveryTestRegion(t, SubscriptionID(1), 1) + region.subscribedSpan.priorityPolicy = newScanPriorityPolicy(pdutil.NewClock4Test(), 30*time.Minute) + + rangeRetryCh := make(chan rangeTask, 2) + handler := newRegionFailureHandler( + regionCache, + func(*subscribedSpan) {}, + func(context.Context, regionInfo) { + t.Fatal("unexpected region retry") + }, + func(_ context.Context, task rangeTask) { + rangeRetryCh <- task + }, + ) + errInfo := newRegionErrorInfo(region, &eventError{ + err: &cdcpb.Error{NotLeader: &errorpb.NotLeader{}}, + }) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + require.NoError(t, handler.handleError(ctx, errInfo)) + + select { + case task := <-rangeRetryCh: + require.Equal(t, region.span, task.span) + require.Same(t, region.subscribedSpan, task.subscribedSpan) + case <-time.After(time.Second): + t.Fatal("not leader retry was not scheduled") + } +} + +func TestRegionRecoveryBackoffFollowsRangeAcrossRegionChanges(t *testing.T) { + regionRetryCh := make(chan regionInfo, 2) + handler := newRegionFailureHandler( + nil, + func(*subscribedSpan) {}, + func(_ context.Context, region regionInfo) { + regionRetryCh <- region + }, + func(context.Context, rangeTask) {}, + ) + t.Cleanup(handler.cancelRecoveries) + + region := createFailureRecoveryTestRegion(t, SubscriptionID(1), 1) + errInfo := newRegionErrorInfo(region, &eventError{ + err: &cdcpb.Error{Congested: &cdcpb.Congested{}}, + }) + require.NoError(t, handler.handleError(context.Background(), errInfo)) + select { + case retried := <-regionRetryCh: + require.Equal(t, uint64(1), retried.verID.GetID()) + case <-time.After(time.Second): + t.Fatal("first region recovery was not scheduled") + } + + region.verID = tikv.NewRegionVerID(2, 1, 1) + errInfo = newRegionErrorInfo(region, &eventError{ + err: &cdcpb.Error{Congested: &cdcpb.Congested{}}, + }) + require.NoError(t, handler.handleError(context.Background(), errInfo)) + select { + case retried := <-regionRetryCh: + require.Equal(t, uint64(2), retried.verID.GetID()) + case <-time.After(time.Second): + t.Fatal("second region recovery was not scheduled") + } + + key := newRegionRecoveryKey(region.subscribedSpan.subID, region.span) + handler.recoveryMu.Lock() + attempt := handler.recoveries[key].attempt + handler.recoveryMu.Unlock() + require.Equal(t, uint32(2), attempt) +} + +func TestRegionFailureHandlerRequestCancelledResetsRecoveryState(t *testing.T) { + handler := newRegionFailureHandler(nil, func(*subscribedSpan) {}, func(context.Context, regionInfo) {}, func(context.Context, rangeTask) {}) + region := createFailureRecoveryTestRegion(t, SubscriptionID(1), 1) + key := newRegionRecoveryKey(region.subscribedSpan.subID, region.span) + handler.recoveries[key] = ®ionRecoveryState{} + + err := handler.handleError(context.Background(), newRegionErrorInfo(region, &requestCancelledErr{})) + require.NoError(t, err) + + handler.recoveryMu.Lock() + _, ok := handler.recoveries[key] + handler.recoveryMu.Unlock() + assert.False(t, ok) +} + +func TestRegionFailureHandlerExpiresRecoveryStates(t *testing.T) { + handler := newRegionFailureHandler(nil, func(*subscribedSpan) {}, func(context.Context, regionInfo) {}, func(context.Context, rangeTask) {}) + now := time.Now() + expiredKey := newRegionRecoveryKey(1, heartbeatpb.TableSpan{StartKey: []byte("a"), EndKey: []byte("b")}) + activeKey := newRegionRecoveryKey(1, heartbeatpb.TableSpan{StartKey: []byte("b"), EndKey: []byte("c")}) + handler.recoveries[expiredKey] = ®ionRecoveryState{expiresAt: now.Add(-time.Second)} + handler.recoveries[activeKey] = ®ionRecoveryState{expiresAt: now.Add(time.Second)} + + handler.expireRecoveries(now) + + handler.recoveryMu.Lock() + _, expiredExists := handler.recoveries[expiredKey] + _, activeExists := handler.recoveries[activeKey] + handler.recoveryMu.Unlock() + require.False(t, expiredExists) + require.True(t, activeExists) +} diff --git a/logservice/logpuller/region_req_cache.go b/logservice/logpuller/region_req_cache.go deleted file mode 100644 index b4478cab07..0000000000 --- a/logservice/logpuller/region_req_cache.go +++ /dev/null @@ -1,341 +0,0 @@ -// Copyright 2025 PingCAP, Inc. -// -// 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, -// See the License for the specific language governing permissions and -// limitations under the License. - -package logpuller - -import ( - "context" - "sync" - "time" - - "github.com/pingcap/log" - "github.com/pingcap/ticdc/pkg/metrics" - "go.uber.org/atomic" - "go.uber.org/zap" -) - -const ( - checkStaleRequestInterval = time.Second * 10 - requestGCLifeTime = time.Minute * 180 - addReqRetryInterval = time.Millisecond * 1 - addReqRetryLimit = 3 - abnormalRequestDurationInSec = 60 * 60 * 2 // 2 hours -) - -// regionReq represents a wrapped region request with state -type regionReq struct { - regionInfo regionInfo - createTime time.Time -} - -func newRegionReq(region regionInfo) regionReq { - return regionReq{ - regionInfo: region, - createTime: time.Now(), - } -} - -func (r *regionReq) isStale() bool { - return time.Since(r.createTime) > requestGCLifeTime -} - -// requestCache manages region requests with flow control -type requestCache struct { - // pending requests waiting to be sent - pendingQueue chan regionReq - - // sent requests waiting for initialization (subscriptionID -> regions -> regionReq) - sentRequests struct { - sync.RWMutex - regionReqs map[SubscriptionID]map[uint64]regionReq - } - - // pendingCount is a flow control slot counter. - // A slot is acquired when a request is successfully enqueued into pendingQueue (see add), - // and is released when the request is finished/removed (resolve/markStopped/markDone/clear). - // pop and markSent don't change it. If markSent overwrites an existing request for the same region, - // it will release a slot for the replaced request to avoid leaking pendingCount. - pendingCount atomic.Int64 - // maximum number of pending requests allowed - maxPendingCount int64 - - // channel to signal when space becomes available - spaceAvailable chan struct{} - - lastCheckStaleRequestTime atomic.Time -} - -func newRequestCache(maxPendingCount int) *requestCache { - res := &requestCache{ - pendingQueue: make(chan regionReq, maxPendingCount), // Large buffer to reduce blocking - sentRequests: struct { - sync.RWMutex - regionReqs map[SubscriptionID]map[uint64]regionReq - }{regionReqs: make(map[SubscriptionID]map[uint64]regionReq)}, - pendingCount: atomic.Int64{}, - maxPendingCount: int64(maxPendingCount), - spaceAvailable: make(chan struct{}, 16), // Buffered to avoid blocking - } - - res.lastCheckStaleRequestTime.Store(time.Now()) - return res -} - -// add adds a new region request to the cache -// It blocks if pendingCount >= maxPendingCount until there's space or ctx is cancelled -func (c *requestCache) add(ctx context.Context, region regionInfo, force bool) (bool, error) { - start := time.Now() - ticker := time.NewTicker(addReqRetryInterval) - defer ticker.Stop() - addReqRetryLimit := addReqRetryLimit - - for { - current := c.pendingCount.Load() - if current < c.maxPendingCount || force { - // Try to add the request - req := newRegionReq(region) - select { - case <-ctx.Done(): - return false, ctx.Err() - case c.pendingQueue <- req: - c.pendingCount.Inc() - cost := time.Since(start) - metrics.SubscriptionClientAddRegionRequestDuration.Observe(cost.Seconds()) - return true, nil - case <-ticker.C: - addReqRetryLimit-- - if addReqRetryLimit <= 0 { - return false, nil - } - continue - } - } - - // Wait for space to become available - select { - case <-ticker.C: - addReqRetryLimit-- - if addReqRetryLimit <= 0 { - return false, nil - } - continue - case <-c.spaceAvailable: - continue - case <-ctx.Done(): - return false, ctx.Err() - } - } -} - -// pop gets the next pending request. -// Note: it doesn't change pendingCount. The slot acquired in add() should be released later -// (e.g. resolve/markStopped/markDone). -func (c *requestCache) pop(ctx context.Context) (regionReq, error) { - select { - case req := <-c.pendingQueue: - return req, nil - case <-ctx.Done(): - return regionReq{}, ctx.Err() - } -} - -// markSent marks a request as sent and adds it to sent requests. -// It doesn't change pendingCount: the slot is released when the request is finished/removed. -func (c *requestCache) markSent(req regionReq) { - c.sentRequests.Lock() - defer c.sentRequests.Unlock() - - m, ok := c.sentRequests.regionReqs[req.regionInfo.subscribedSpan.subID] - - if !ok { - m = make(map[uint64]regionReq) - c.sentRequests.regionReqs[req.regionInfo.subscribedSpan.subID] = m - } - - if oldReq, exists := m[req.regionInfo.verID.GetID()]; exists { - log.Warn("region request overwritten", - zap.Uint64("subID", uint64(req.regionInfo.subscribedSpan.subID)), - zap.Uint64("regionID", req.regionInfo.verID.GetID()), - zap.Float64("oldAgeSec", time.Since(oldReq.createTime).Seconds()), - zap.Float64("newAgeSec", time.Since(req.createTime).Seconds()), - zap.Int("pendingCount", int(c.pendingCount.Load())), - zap.Int("pendingQueueLen", len(c.pendingQueue))) - c.markDone() - } - m[req.regionInfo.verID.GetID()] = req -} - -// markStopped removes a sent request and releases a slot. -func (c *requestCache) markStopped(subID SubscriptionID, regionID uint64) { - c.sentRequests.Lock() - defer c.sentRequests.Unlock() - - regionReqs, ok := c.sentRequests.regionReqs[subID] - if !ok { - return - } - - _, exists := regionReqs[regionID] - if !exists { - return - } - - delete(regionReqs, regionID) - if len(regionReqs) == 0 { - delete(c.sentRequests.regionReqs, subID) - } - c.markDone() -} - -// resolve marks a region as initialized and removes it from sent requests -func (c *requestCache) resolve(subscriptionID SubscriptionID, regionID uint64) bool { - c.sentRequests.Lock() - defer c.sentRequests.Unlock() - regionReqs, ok := c.sentRequests.regionReqs[subscriptionID] - if !ok { - return false - } - - req, exists := regionReqs[regionID] - if !exists { - return false - } - - // Check if the subscription ID matches - if req.regionInfo.subscribedSpan.subID == subscriptionID { - delete(regionReqs, regionID) - c.markDone() - cost := time.Since(req.createTime).Seconds() - if cost > 0 && cost < abnormalRequestDurationInSec { - log.Debug("cdc resolve region request", zap.Uint64("subID", uint64(subscriptionID)), zap.Uint64("regionID", regionID), zap.Float64("cost", cost), zap.Int("pendingCount", int(c.pendingCount.Load())), zap.Int("pendingQueueLen", len(c.pendingQueue))) - metrics.RegionRequestFinishScanDuration.Observe(cost) - } else { - log.Info("region request duration abnormal, skip metric", zap.Float64("cost", cost), zap.Uint64("regionID", regionID)) - } - return true - } - - return false -} - -// clearStaleRequest clears stale requests from the cache -// Note: Sometimes, the CDC sends the same region request to TiKV multiple times. In such cases, this method is needed to reduce the pendingSize. -func (c *requestCache) clearStaleRequest() { - if time.Since(c.lastCheckStaleRequestTime.Load()) < checkStaleRequestInterval { - return - } - c.sentRequests.Lock() - defer c.sentRequests.Unlock() - reqCount := 0 - for subID, regionReqs := range c.sentRequests.regionReqs { - for regionID, regionReq := range regionReqs { - if regionReq.regionInfo.isStopped() || - regionReq.regionInfo.subscribedSpan.stopped.Load() || - regionReq.regionInfo.lockedRangeState.Initialized.Load() || - regionReq.isStale() { - c.markDone() - log.Warn("region worker delete stale region request", - zap.Uint64("subID", uint64(subID)), - zap.Uint64("regionID", regionID), - zap.Int("pendingCount", int(c.pendingCount.Load())), - zap.Int("pendingQueueLen", len(c.pendingQueue)), - zap.Bool("isRegionStopped", regionReq.regionInfo.isStopped()), - zap.Bool("isSubscribedSpanStopped", regionReq.regionInfo.subscribedSpan.stopped.Load()), - zap.Bool("isStale", regionReq.isStale()), - zap.Time("createTime", regionReq.createTime)) - delete(regionReqs, regionID) - } else { - reqCount++ - } - } - if len(regionReqs) == 0 { - delete(c.sentRequests.regionReqs, subID) - } - } - - // If there are no in-cache region requests but pendingCount isn't 0, it means pendingCount is stale. - // Reset it to avoid blocking add() forever. - if reqCount == 0 && len(c.pendingQueue) == 0 && c.pendingCount.Load() != 0 { - log.Info("region worker pending request count is not equal to actual region request count, correct it", - zap.Int("pendingCount", int(c.pendingCount.Load())), - zap.Int("actualReqCount", reqCount), - zap.Int("pendingQueueLen", len(c.pendingQueue))) - c.pendingCount.Store(0) - // Notify waiting add operations that there's space available. - select { - case c.spaceAvailable <- struct{}{}: - default: - } - } - - c.lastCheckStaleRequestTime.Store(time.Now()) -} - -// clear removes all requests and returns them -func (c *requestCache) clear() []regionInfo { - var regions []regionInfo - - // Drain pending requests from channel -LOOP: - for { - select { - case req := <-c.pendingQueue: - regions = append(regions, req.regionInfo) - c.markDone() - default: - break LOOP - } - } - - c.sentRequests.Lock() - defer c.sentRequests.Unlock() - - for subID, regionReqs := range c.sentRequests.regionReqs { - for regionID := range regionReqs { - regions = append(regions, regionReqs[regionID].regionInfo) - delete(regionReqs, regionID) - c.markDone() - } - delete(c.sentRequests.regionReqs, subID) - } - return regions -} - -// getPendingCount returns the current pending count -func (c *requestCache) getPendingCount() int { - return int(c.pendingCount.Load()) -} - -func (c *requestCache) markDone() { - // Decrement pendingCount by 1, but never let it go below 0. - // Do it with CAS to avoid clobbering concurrent Inc() calls. - for { - old := c.pendingCount.Load() - if old == 0 { - break - } else if old < 0 { - if c.pendingCount.CompareAndSwap(old, 0) { - break - } - } else { - if c.pendingCount.CompareAndSwap(old, old-1) { - break - } - } - } - // Notify waiting add operations that there's space available. - select { - case c.spaceAvailable <- struct{}{}: - default: // If channel is full, skip notification - } -} diff --git a/logservice/logpuller/region_req_cache_test.go b/logservice/logpuller/region_req_cache_test.go deleted file mode 100644 index 62706a8542..0000000000 --- a/logservice/logpuller/region_req_cache_test.go +++ /dev/null @@ -1,336 +0,0 @@ -// Copyright 2025 PingCAP, Inc. -// -// 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, -// See the License for the specific language governing permissions and -// limitations under the License. - -package logpuller - -import ( - "context" - "testing" - "time" - - "github.com/pingcap/ticdc/heartbeatpb" - "github.com/stretchr/testify/require" - "github.com/tikv/client-go/v2/tikv" -) - -func createTestRegionInfo(subID SubscriptionID, regionID uint64) regionInfo { - verID := tikv.NewRegionVerID(regionID, 1, 1) - - span := heartbeatpb.TableSpan{ - TableID: 1, - StartKey: []byte("start"), - EndKey: []byte("end"), - } - - subscribedSpan := &subscribedSpan{ - subID: subID, - startTs: 100, - span: span, - } - - return newRegionInfo(verID, span, nil, subscribedSpan, false) -} - -func TestRequestCacheAdd_NormalCase(t *testing.T) { - cache := newRequestCache(10) - ctx := context.Background() - - region := createTestRegionInfo(1, 1) - - ok, err := cache.add(ctx, region, false) - require.NoError(t, err) - require.True(t, ok) - require.Equal(t, 1, cache.getPendingCount()) - - // Verify the request was added to the queue - req, err := cache.pop(ctx) - require.NoError(t, err) - require.NotNil(t, req) - require.Equal(t, region.verID.GetID(), req.regionInfo.verID.GetID()) - require.Equal(t, region.subscribedSpan.subID, req.regionInfo.subscribedSpan.subID) -} - -func TestRequestCacheAdd_ForceFlag(t *testing.T) { - cache := newRequestCache(1) - ctx := context.Background() - - // Fill up the cache - region1 := createTestRegionInfo(1, 1) - ok, err := cache.add(ctx, region1, false) - require.True(t, ok) - require.NoError(t, err) - require.Equal(t, 1, cache.getPendingCount()) - - // Try to add another request without force - should fail due to retry limit - region2 := createTestRegionInfo(1, 2) - ok, err = cache.add(ctx, region2, false) - require.False(t, ok) - require.NoError(t, err) - - // With force=true, it should still fail because the channel is full - // The force flag only bypasses the pendingCount check, not the channel capacity - region3 := createTestRegionInfo(1, 3) - ok, err = cache.add(ctx, region3, true) - require.False(t, ok) - require.NoError(t, err) - - // consume the pending queue ann add with force - req, err := cache.pop(ctx) - require.NoError(t, err) - require.NotNil(t, req) - require.Equal(t, region1.verID.GetID(), req.regionInfo.verID.GetID()) - require.Equal(t, region1.subscribedSpan.subID, req.regionInfo.subscribedSpan.subID) - cache.markSent(req) - require.Equal(t, 1, cache.getPendingCount()) - - ok, err = cache.add(ctx, region3, true) - require.True(t, ok) - require.NoError(t, err) - // It is 2 since region1 is unresolved - require.Equal(t, 2, cache.getPendingCount()) - - // resolve region1 - cache.resolve(region1.subscribedSpan.subID, region1.verID.GetID()) - require.Equal(t, 1, cache.getPendingCount()) -} - -func TestRequestCacheAdd_ContextCancellation(t *testing.T) { - cache := newRequestCache(1) - - // Fill up the cache - region1 := createTestRegionInfo(1, 1) - ctx1 := context.Background() - ok, err := cache.add(ctx1, region1, false) - require.True(t, ok) - require.NoError(t, err) - - // Try to add another request with a cancelled context - ctx2, cancel := context.WithCancel(context.Background()) - cancel() // Cancel immediately - - region2 := createTestRegionInfo(1, 2) - ok, err = cache.add(ctx2, region2, false) - require.False(t, ok) - require.Error(t, err) - require.Equal(t, context.Canceled, err) -} - -func TestRequestCacheAdd_RetryLimitExceeded(t *testing.T) { - cache := newRequestCache(1) - ctx := context.Background() - - // Fill up the cache - region1 := createTestRegionInfo(1, 1) - ok, err := cache.add(ctx, region1, false) - require.True(t, ok) - require.NoError(t, err) - - // Try to add another request - should eventually hit retry limit - region2 := createTestRegionInfo(1, 2) - ok, err = cache.add(ctx, region2, false) - require.False(t, ok) - require.NoError(t, err) -} - -func TestRequestCacheAdd_SpaceAvailableNotification(t *testing.T) { - cache := newRequestCache(2) - ctx := context.Background() - - // Fill up the cache - region1 := createTestRegionInfo(1, 1) - ok, err := cache.add(ctx, region1, false) - require.True(t, ok) - require.NoError(t, err) - require.Equal(t, 1, cache.getPendingCount()) - - region2 := createTestRegionInfo(1, 2) - ok, err = cache.add(ctx, region2, false) - require.True(t, ok) - require.NoError(t, err) - require.Equal(t, 2, cache.getPendingCount()) - - // Pop a request and mark it as sent, then resolve it to free up space - req, err := cache.pop(ctx) - require.NoError(t, err) - require.NotNil(t, req) - require.Equal(t, 2, cache.getPendingCount()) // pop doesn't change pendingCount - // Mark as sent - cache.markSent(req) - require.Equal(t, 2, cache.getPendingCount()) - - // Resolve the request to free up space - success := cache.resolve(req.regionInfo.subscribedSpan.subID, req.regionInfo.verID.GetID()) - require.True(t, success) - require.Equal(t, 1, cache.getPendingCount()) - - // Now we should be able to add another request - region3 := createTestRegionInfo(1, 3) - ok, err = cache.add(ctx, region3, false) - require.True(t, ok) - require.NoError(t, err) - require.Equal(t, 2, cache.getPendingCount()) -} - -func TestRequestCacheAdd_ConcurrentAdds(t *testing.T) { - cache := newRequestCache(10) - ctx := context.Background() - - const numGoroutines = 5 - done := make(chan error, numGoroutines) - - // Start multiple goroutines adding requests concurrently - for i := 0; i < numGoroutines; i++ { - go func(id int) { - region := createTestRegionInfo(SubscriptionID(id%3), uint64(id)) - ok, err := cache.add(ctx, region, false) - require.True(t, ok) - require.NoError(t, err) - done <- err - }(i) - } - - // Wait for all goroutines to complete - for i := 0; i < numGoroutines; i++ { - select { - case err := <-done: - require.NoError(t, err) - case <-time.After(1 * time.Second): - t.Fatal("Timeout waiting for concurrent adds to complete") - } - } - - require.Equal(t, numGoroutines, cache.getPendingCount()) -} - -func TestRequestCacheAdd_StaleRequestCleanup(t *testing.T) { - cache := newRequestCache(10) - ctx := context.Background() - - // Add a request and mark it as sent - region := createTestRegionInfo(1, 1) - ok, err := cache.add(ctx, region, false) - require.True(t, ok) - require.NoError(t, err) - - req, err := cache.pop(ctx) - require.NoError(t, err) - require.NotNil(t, req) - - // Mark as sent - cache.markSent(req) - require.Equal(t, 1, cache.getPendingCount()) - - // Manually set the request as stale by modifying createTime - cache.sentRequests.Lock() - regionReqs := cache.sentRequests.regionReqs[req.regionInfo.subscribedSpan.subID] - regionReqs[req.regionInfo.verID.GetID()] = regionReq{ - regionInfo: req.regionInfo, - createTime: time.Now().Add(-requestGCLifeTime - time.Second), // Make it stale - } - cache.sentRequests.Unlock() - - // Manually set lastCheckStaleRequestTime to bypass the time interval check - cache.lastCheckStaleRequestTime.Store(time.Now().Add(-checkStaleRequestInterval - time.Second)) - - // Manually trigger stale cleanup by calling clearStaleRequest - cache.clearStaleRequest() - - // The stale request should be cleaned up - require.Equal(t, 0, cache.getPendingCount()) -} - -func TestRequestCacheAdd_WithStoppedRegion(t *testing.T) { - cache := newRequestCache(10) - ctx := context.Background() - - // Create a region info with stopped state (lockedRangeState = nil) - region := createTestRegionInfo(1, 1) - region.lockedRangeState = nil // This makes it stopped - - ok, err := cache.add(ctx, region, false) - require.True(t, ok) - require.NoError(t, err) - require.Equal(t, 1, cache.getPendingCount()) - - req, err := cache.pop(ctx) - require.NoError(t, err) - require.NotNil(t, req) - - // Mark as sent - cache.markSent(req) - require.Equal(t, 1, cache.getPendingCount()) - - // Manually set lastCheckStaleRequestTime to bypass the time interval check - cache.lastCheckStaleRequestTime.Store(time.Now().Add(-checkStaleRequestInterval - time.Second)) - - // Manually trigger cleanup of stopped region - cache.clearStaleRequest() - - // The stopped region should be cleaned up - require.Equal(t, 0, cache.getPendingCount()) -} - -func TestRequestCacheMarkSent_DuplicateReleaseSlot(t *testing.T) { - cache := newRequestCache(10) - ctx := context.Background() - - region := createTestRegionInfo(1, 1) - - ok, err := cache.add(ctx, region, false) - require.True(t, ok) - require.NoError(t, err) - - // Add a duplicate request for the same region. It should not leak pendingCount even if - // markSent overwrites the existing entry. - ok, err = cache.add(ctx, region, false) - require.True(t, ok) - require.NoError(t, err) - require.Equal(t, 2, cache.getPendingCount()) - - req1, err := cache.pop(ctx) - require.NoError(t, err) - cache.markSent(req1) - require.Equal(t, 2, cache.getPendingCount()) - - req2, err := cache.pop(ctx) - require.NoError(t, err) - cache.markSent(req2) - require.Equal(t, 1, cache.getPendingCount()) - - // Finish the remaining tracked request. - require.True(t, cache.resolve(region.subscribedSpan.subID, region.verID.GetID())) - require.Equal(t, 0, cache.getPendingCount()) -} - -func TestRequestCacheMarkStopped_ReleasesSlot(t *testing.T) { - cache := newRequestCache(10) - ctx := context.Background() - - region := createTestRegionInfo(1, 1) - - ok, err := cache.add(ctx, region, false) - require.True(t, ok) - require.NoError(t, err) - require.Equal(t, 1, cache.getPendingCount()) - - req, err := cache.pop(ctx) - require.NoError(t, err) - - cache.markSent(req) - require.Equal(t, 1, cache.getPendingCount()) - require.Contains(t, cache.sentRequests.regionReqs, req.regionInfo.subscribedSpan.subID) - - cache.markStopped(req.regionInfo.subscribedSpan.subID, req.regionInfo.verID.GetID()) - require.Equal(t, 0, cache.getPendingCount()) - require.NotContains(t, cache.sentRequests.regionReqs, req.regionInfo.subscribedSpan.subID) -} diff --git a/logservice/logpuller/region_request_scheduler.go b/logservice/logpuller/region_request_scheduler.go new file mode 100644 index 0000000000..d71d5af62a --- /dev/null +++ b/logservice/logpuller/region_request_scheduler.go @@ -0,0 +1,223 @@ +// Copyright 2026 PingCAP, Inc. +// +// 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 logpuller + +import ( + "context" + "sync" + "sync/atomic" + + "github.com/pingcap/log" + "github.com/pingcap/ticdc/pkg/common" + "github.com/pingcap/ticdc/pkg/config" + "github.com/pingcap/ticdc/pkg/errors" + "github.com/pingcap/ticdc/pkg/metrics" + "github.com/pingcap/ticdc/utils/priorityqueue" + kvclientv2 "github.com/tikv/client-go/v2/kv" + "github.com/tikv/client-go/v2/tikv" + "go.uber.org/zap" + "go.uber.org/zap/zapcore" + "golang.org/x/sync/errgroup" +) + +const regionRequestWorkerPerStore = 8 + +// regionRequestScheduler routes locked region requests through the global +// priority queue to a worker connected to the region's TiKV store. Range +// resolution and retry policy remain owned by subscriptionClient and +// regionFailureHandler respectively. +type regionRequestScheduler struct { + upstream *upstreamHandle + eventSink *regionEventSink + failureHandler *regionFailureHandler + memoryQuota *memoryQuotaController + + // taskQueue orders all regions before they are assigned to a TiKV store. + taskQueue *priorityqueue.PriorityQueue[*regionPriorityTask] + // sequence is the FIFO tie-breaker for regions in the same priority class. + sequence atomic.Uint64 + // stores maps TiKV addresses to regionRequestStore. Stores are created only + // by Run, but are also read by metrics and deregistration goroutines. + stores sync.Map + + // workerCount is the configured number of request workers per store. + workerCount int + // workerWindow is each worker's share of the configured store window. + workerWindow int + // maxWindowMultiplier is passed to each worker's admission controller. + maxWindowMultiplier int +} + +func newRegionRequestScheduler( + upstream *upstreamHandle, + eventSink *regionEventSink, + failureHandler *regionFailureHandler, + memoryQuota *memoryQuotaController, +) *regionRequestScheduler { + pullerConfig := config.GetGlobalServerConfig().Debug.Puller + workerCount := regionRequestWorkerPerStore + workerWindow := (pullerConfig.PendingRegionRequestQueueSize + workerCount - 1) / workerCount + return ®ionRequestScheduler{ + upstream: upstream, + eventSink: eventSink, + failureHandler: failureHandler, + memoryQuota: memoryQuota, + taskQueue: priorityqueue.New[*regionPriorityTask](), + workerCount: workerCount, + workerWindow: workerWindow, + maxWindowMultiplier: pullerConfig.RegionRequestMaxWindowMultiplier, + } +} + +func (s *regionRequestScheduler) Submit(region regionInfo) { + if log.GetLevel() <= zapcore.DebugLevel { + log.Debug("cdc region scan task enqueued", + zap.Uint64("subscriptionID", uint64(region.subscribedSpan.subID)), + zap.Int64("tableID", region.subscribedSpan.span.TableID), + zap.Uint64("startTs", region.subscribedSpan.startTs), + zap.Uint64("regionID", region.verID.GetID()), + zap.Uint64("regionEpochVersion", region.verID.GetVer()), + zap.Uint64("regionEpochConfVer", region.verID.GetConfVer()), + zap.String("priority", normalizeScanPriority(region.scanPriority).String()), + zap.String("scanPriority", region.scanPriority.String()), + zap.String("span", common.FormatTableSpan(®ion.span))) + } + s.taskQueue.Push(newRegionPriorityTask(region, s.sequence.Add(1))) +} + +func (s *regionRequestScheduler) Run(ctx context.Context, workerGroup *errgroup.Group) error { + defer func() { + s.stores.Range(func(_, value any) bool { + value.(*regionRequestStore).close() + return true + }) + }() + + for { + select { + case <-ctx.Done(): + return ctx.Err() + default: + } + + task, err := s.taskQueue.Pop(ctx) + if err != nil { + if errors.Is(err, priorityqueue.ErrClosed) { + return nil + } + return err + } + + region, err := s.attachRPCContext(ctx, task.regionInfo) + if err != nil { + s.failureHandler.Report(newRegionErrorInfo(region, err)) + continue + } + if region.subscribedSpan.stopped.Load() { + s.failureHandler.Report(newRegionErrorInfo(region, &requestCancelledErr{})) + continue + } + + store := s.getOrCreateStore(ctx, workerGroup, region.rpcCtx.Addr) + task.regionInfo = region + if !store.submit(task) { + if ctx.Err() != nil { + return ctx.Err() + } + s.failureHandler.Report(newRegionErrorInfo(region, &storeStreamErr{})) + continue + } + if log.GetLevel() <= zapcore.DebugLevel { + log.Debug("subscription client will request a region", + zap.Uint64("subscriptionID", uint64(region.subscribedSpan.subID)), + zap.Uint64("regionID", region.verID.GetID()), + zap.String("addr", region.rpcCtx.Addr)) + } + } +} + +func (s *regionRequestScheduler) attachRPCContext( + ctx context.Context, + region regionInfo, +) (regionInfo, error) { + bo := tikv.NewBackoffer(ctx, tikvRequestMaxBackoff) + rpcCtx, err := s.upstream.regionCache.GetTiKVRPCContext( + bo, region.verID, kvclientv2.ReplicaReadLeader, 0) + if rpcCtx != nil { + region.rpcCtx = rpcCtx + return region, nil + } + if err != nil { + log.Debug("region request scheduler failed to get RPC context", + zap.Uint64("subscriptionID", uint64(region.subscribedSpan.subID)), + zap.Uint64("regionID", region.verID.GetID()), + zap.Error(err)) + } + return region, &rpcCtxUnavailableErr{verID: region.verID} +} + +func (s *regionRequestScheduler) getOrCreateStore( + ctx context.Context, + workerGroup *errgroup.Group, + storeAddr string, +) *regionRequestStore { + if value, ok := s.stores.Load(storeAddr); ok { + return value.(*regionRequestStore) + } + + store := newRegionRequestStore( + s.upstream, + s.eventSink, + s.failureHandler, + storeAddr, + s.workerCount, + s.workerWindow, + s.maxWindowMultiplier, + s.memoryQuota, + ) + // The scheduler run loop is the only writer. Publish the store after its + // immutable worker list is complete, then start its workers. + s.stores.Store(storeAddr, store) + store.startWorkers(ctx, workerGroup) + return store +} + +func (s *regionRequestScheduler) BroadcastDeregister( + subID SubscriptionID, + filterLoop bool, +) { + s.stores.Range(func(_, value any) bool { + value.(*regionRequestStore).broadcastDeregister(subID, filterLoop) + return true + }) +} + +func (s *regionRequestScheduler) inflightCount() int { + count := 0 + s.stores.Range(func(_, value any) bool { + count += value.(*regionRequestStore).inflightCount() + return true + }) + return count +} + +func (s *regionRequestScheduler) UpdateMetrics() { + metrics.SubscriptionClientRequestedRegionCount.WithLabelValues("inflight"). + Set(float64(s.inflightCount())) +} + +func (s *regionRequestScheduler) Close() { + s.taskQueue.Close() +} diff --git a/logservice/logpuller/region_request_scheduler_test.go b/logservice/logpuller/region_request_scheduler_test.go new file mode 100644 index 0000000000..6180d362df --- /dev/null +++ b/logservice/logpuller/region_request_scheduler_test.go @@ -0,0 +1,247 @@ +// Copyright 2026 PingCAP, Inc. +// +// 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 logpuller + +import ( + "context" + "testing" + "time" + + "github.com/pingcap/ticdc/heartbeatpb" + "github.com/pingcap/ticdc/logservice/logpuller/regionlock" + "github.com/pingcap/ticdc/pkg/pdutil" + "github.com/pingcap/ticdc/utils/priorityqueue" + "github.com/pingcap/tidb/pkg/store/mockstore/mockcopr" + "github.com/stretchr/testify/require" + "github.com/tikv/client-go/v2/testutils" + "github.com/tikv/client-go/v2/tikv" + "golang.org/x/sync/errgroup" +) + +func TestRegionRequestSchedulerBroadcastDeregisterUsesWorkerControlQueue(t *testing.T) { + scheduler := ®ionRequestScheduler{} + + worker1 := ®ionRequestWorker{ + storeAddr: "store-1", + admission: newTestRegionAdmissionController(1, 1), + controlQueue: newControlQueue(), + } + worker2 := ®ionRequestWorker{ + storeAddr: "store-2", + admission: newTestRegionAdmissionController(1, 1), + controlQueue: newControlQueue(), + } + store1 := ®ionRequestStore{workers: []*regionRequestWorker{worker1}} + store2 := ®ionRequestStore{workers: []*regionRequestWorker{worker2}} + scheduler.stores.Store("store-1", store1) + scheduler.stores.Store("store-2", store2) + + dummyRegion := regionInfo{ + subscribedSpan: &subscribedSpan{subID: SubscriptionID(2)}, + lockedRangeState: ®ionlock.LockedRangeState{}, + } + require.True(t, worker1.admission.submit(newRegionPriorityTask(dummyRegion, 1))) + + scheduler.BroadcastDeregister(SubscriptionID(1), true) + + req1, ok := worker1.controlQueue.tryPop() + require.True(t, ok) + require.Equal(t, SubscriptionID(1), req1.subID) + require.True(t, req1.filterLoop) + + req2, ok := worker2.controlQueue.tryPop() + require.True(t, ok) + require.Equal(t, SubscriptionID(1), req2.subID) + require.True(t, req2.filterLoop) + + require.Equal(t, 1, worker1.admission.stats().pending) +} + +func TestRegionRequestSchedulerInflightCountAggregatesStores(t *testing.T) { + scheduler := ®ionRequestScheduler{} + + worker1 := ®ionRequestWorker{admission: newTestRegionAdmissionController(2, 1)} + worker2 := ®ionRequestWorker{admission: newTestRegionAdmissionController(2, 1)} + scheduler.stores.Store("store-1", ®ionRequestStore{workers: []*regionRequestWorker{worker1}}) + scheduler.stores.Store("store-2", ®ionRequestStore{workers: []*regionRequestWorker{worker2}}) + + req1 := admitRegionRequest(t, worker1.admission, prepareRegionForAdmission(createTestRegionInfo(1, 1), 100)) + req2 := admitRegionRequest(t, worker2.admission, prepareRegionForAdmission(createTestRegionInfo(1, 2), 100)) + + require.Equal(t, 2, scheduler.inflightCount()) + + require.True(t, req1.abort()) + require.True(t, req2.abort()) +} + +func TestRegionRequestSchedulerReschedulesRegionWhenStoreSubmitFails(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + _, cluster, pdClient, _ := testutils.NewMockTiKV("", mockcopr.NewCoprRPCHandler()) + pdClient = &mockPDClient{Client: pdClient, versionGen: defaultVersionGen} + defer pdClient.Close() + + const storeAddr = "store-1" + cluster.AddStore(1, storeAddr) + cluster.Bootstrap(11, []uint64{1}, []uint64{2}, 2) + + regionCache := tikv.NewRegionCache(pdClient) + defer regionCache.Close() + + bo := tikv.NewBackoffer(ctx, tikvRequestMaxBackoff) + location, err := regionCache.LocateKey(bo, []byte("a")) + require.NoError(t, err) + + rawSpan := heartbeatpb.TableSpan{ + TableID: 1, + StartKey: []byte("a"), + EndKey: []byte("b"), + } + span := &subscribedSpan{ + subID: SubscriptionID(1), + startTs: 100, + span: rawSpan, + rangeLock: regionlock.NewRangeLock(1, rawSpan.StartKey, rawSpan.EndKey, 100), + priorityPolicy: newScanPriorityPolicy(pdutil.NewClock4Test(), 30*time.Minute), + } + lockRes := span.rangeLock.LockRange( + context.Background(), rawSpan.StartKey, rawSpan.EndKey, location.Region.GetID(), location.Region.GetVer()) + require.Equal(t, regionlock.LockRangeStatusSuccess, lockRes.Status) + + admission := newTestRegionAdmissionController(1, 1) + admission.close() + store := ®ionRequestStore{workers: []*regionRequestWorker{{admission: admission}}} + + handler := newRegionFailureHandler(nil, func(*subscribedSpan) {}, nil, nil) + scheduler := ®ionRequestScheduler{ + upstream: &upstreamHandle{ + pd: pdClient, + regionCache: regionCache, + }, + failureHandler: handler, + taskQueue: priorityqueue.New[*regionPriorityTask](), + } + scheduler.stores.Store(storeAddr, store) + + region := newRegionInfo(location.Region, rawSpan, nil, span, false) + region.lockedRangeState = lockRes.LockedRangeState + scheduler.taskQueue.Push(newRegionPriorityTask(region, 1)) + + var workerGroup errgroup.Group + errCh := make(chan error, 1) + go func() { + errCh <- scheduler.Run(ctx, &workerGroup) + }() + + require.Eventually(t, func() bool { + return errCacheLen(handler) == 1 + }, time.Second, 20*time.Millisecond) + + select { + case err := <-errCh: + t.Fatalf("scheduler exited unexpectedly: %v", err) + default: + } + + batch := handler.cache.popBatch(1) + require.Len(t, batch, 1) + require.Equal(t, region.verID, batch[0].verID) + var streamErr *storeStreamErr + require.ErrorAs(t, batch[0].err, &streamErr) + + cancel() + require.ErrorIs(t, <-errCh, context.Canceled) +} + +func TestRegionRequestSchedulerSkipsStoppedSubscriptionBeforeCreatingStore(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + _, cluster, pdClient, _ := testutils.NewMockTiKV("", mockcopr.NewCoprRPCHandler()) + pdClient = &mockPDClient{Client: pdClient, versionGen: defaultVersionGen} + defer pdClient.Close() + + const storeAddr = "store-1" + cluster.AddStore(1, storeAddr) + cluster.Bootstrap(11, []uint64{1}, []uint64{2}, 2) + + regionCache := tikv.NewRegionCache(pdClient) + defer regionCache.Close() + + bo := tikv.NewBackoffer(ctx, tikvRequestMaxBackoff) + location, err := regionCache.LocateKey(bo, []byte("a")) + require.NoError(t, err) + + rawSpan := heartbeatpb.TableSpan{ + TableID: 1, + StartKey: []byte("a"), + EndKey: []byte("b"), + } + span := &subscribedSpan{ + subID: SubscriptionID(1), + startTs: 100, + span: rawSpan, + rangeLock: regionlock.NewRangeLock(1, rawSpan.StartKey, rawSpan.EndKey, 100), + priorityPolicy: newScanPriorityPolicy(pdutil.NewClock4Test(), 30*time.Minute), + } + lockRes := span.rangeLock.LockRange( + context.Background(), rawSpan.StartKey, rawSpan.EndKey, location.Region.GetID(), location.Region.GetVer()) + require.Equal(t, regionlock.LockRangeStatusSuccess, lockRes.Status) + span.stopped.Store(true) + require.False(t, span.rangeLock.Stop()) + + drainedCh := make(chan *subscribedSpan, 1) + handler := newRegionFailureHandler(nil, func(rt *subscribedSpan) { + drainedCh <- rt + }, nil, nil) + scheduler := ®ionRequestScheduler{ + upstream: &upstreamHandle{ + pd: pdClient, + regionCache: regionCache, + }, + failureHandler: handler, + taskQueue: priorityqueue.New[*regionPriorityTask](), + } + + region := newRegionInfo(location.Region, rawSpan, nil, span, false) + region.lockedRangeState = lockRes.LockedRangeState + scheduler.taskQueue.Push(newRegionPriorityTask(region, 1)) + + var workerGroup errgroup.Group + errCh := make(chan error, 1) + go func() { + errCh <- scheduler.Run(ctx, &workerGroup) + }() + + select { + case drained := <-drainedCh: + require.Same(t, span, drained) + case <-time.After(time.Second): + t.Fatal("stopped subscription was not drained") + } + + _, ok := scheduler.stores.Load(storeAddr) + require.False(t, ok) + + select { + case err := <-errCh: + t.Fatalf("scheduler exited unexpectedly: %v", err) + default: + } + + cancel() + require.ErrorIs(t, <-errCh, context.Canceled) +} diff --git a/logservice/logpuller/region_request_store.go b/logservice/logpuller/region_request_store.go new file mode 100644 index 0000000000..38393b3acf --- /dev/null +++ b/logservice/logpuller/region_request_store.go @@ -0,0 +1,88 @@ +// Copyright 2026 PingCAP, Inc. +// +// 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 logpuller + +import ( + "context" + "sync/atomic" + + "golang.org/x/sync/errgroup" +) + +// regionRequestStore owns the region request workers connected to one TiKV +// store. The worker slice is complete before the store is published and is +// immutable afterwards, so task submission only needs an atomic round-robin counter. +type regionRequestStore struct { + workers []*regionRequestWorker + nextWorker atomic.Uint64 +} + +func newRegionRequestStore( + upstream *upstreamHandle, + eventSink *regionEventSink, + failureHandler *regionFailureHandler, + storeAddr string, + workerCount int, + workerWindow int, + maxWindowMultiplier int, + memoryQuota *memoryQuotaController, +) *regionRequestStore { + store := ®ionRequestStore{ + workers: make([]*regionRequestWorker, 0, workerCount), + } + for i := 0; i < workerCount; i++ { + store.workers = append(store.workers, newRegionRequestWorker( + upstream, + eventSink, + failureHandler, + storeAddr, + workerWindow, + maxWindowMultiplier, + memoryQuota, + )) + } + return store +} + +func (s *regionRequestStore) startWorkers(ctx context.Context, workerGroup *errgroup.Group) { + for _, worker := range s.workers { + workerGroup.Go(func() error { return worker.Run(ctx) }) + } +} + +func (s *regionRequestStore) submit(task *regionPriorityTask) bool { + index := (s.nextWorker.Add(1) - 1) % uint64(len(s.workers)) + return s.workers[index].admission.submit(task) +} + +func (s *regionRequestStore) broadcastDeregister(subID SubscriptionID, filterLoop bool) { + for _, worker := range s.workers { + worker.controlQueue.push(deregisterRequest{subID: subID, filterLoop: filterLoop}) + } +} + +func (s *regionRequestStore) close() { + for _, worker := range s.workers { + worker.admission.close() + } +} + +func (s *regionRequestStore) inflightCount() int { + count := 0 + for _, worker := range s.workers { + count += worker.admission.stats().inflight + } + return count +} diff --git a/logservice/logpuller/region_request_store_test.go b/logservice/logpuller/region_request_store_test.go new file mode 100644 index 0000000000..a433d47d23 --- /dev/null +++ b/logservice/logpuller/region_request_store_test.go @@ -0,0 +1,70 @@ +// Copyright 2026 PingCAP, Inc. +// +// 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 logpuller + +import ( + "testing" + + "github.com/stretchr/testify/require" + "github.com/tikv/client-go/v2/tikv" +) + +func TestRegionRequestStoreDistributesRegionsAcrossWorkers(t *testing.T) { + worker1 := ®ionRequestWorker{admission: newTestRegionAdmissionController(1, 1)} + worker2 := ®ionRequestWorker{admission: newTestRegionAdmissionController(1, 1)} + store := ®ionRequestStore{ + workers: []*regionRequestWorker{worker1, worker2}, + } + + for i := uint64(1); i <= 4; i++ { + region := prepareRegionForAdmission(createTestRegionInfo(1, i), 100) + region.verID = tikv.NewRegionVerID(i, 1, 1) + require.True(t, store.submit(newRegionPriorityTask(region, i))) + } + + require.Equal(t, 2, worker1.admission.stats().pending) + require.Equal(t, 2, worker2.admission.stats().pending) +} + +func TestRegionRequestStoreInflightCountAggregatesWorkers(t *testing.T) { + worker1 := ®ionRequestWorker{admission: newTestRegionAdmissionController(2, 1)} + worker2 := ®ionRequestWorker{admission: newTestRegionAdmissionController(2, 1)} + store := ®ionRequestStore{ + workers: []*regionRequestWorker{worker1, worker2}, + } + + req1 := admitRegionRequest(t, worker1.admission, prepareRegionForAdmission(createTestRegionInfo(1, 1), 100)) + req2 := admitRegionRequest(t, worker2.admission, prepareRegionForAdmission(createTestRegionInfo(1, 2), 100)) + + require.Equal(t, 2, store.inflightCount()) + + require.True(t, req1.abort()) + require.True(t, req2.abort()) +} + +func TestRegionRequestStoreCloseClosesWorkerAdmissions(t *testing.T) { + worker1 := ®ionRequestWorker{admission: newTestRegionAdmissionController(1, 1)} + worker2 := ®ionRequestWorker{admission: newTestRegionAdmissionController(1, 1)} + store := ®ionRequestStore{ + workers: []*regionRequestWorker{worker1, worker2}, + } + + store.close() + + region1 := prepareRegionForAdmission(createTestRegionInfo(1, 1), 100) + region2 := prepareRegionForAdmission(createTestRegionInfo(1, 2), 100) + require.False(t, worker1.admission.submit(newRegionPriorityTask(region1, 1))) + require.False(t, worker2.admission.submit(newRegionPriorityTask(region2, 2))) +} diff --git a/logservice/logpuller/region_request_worker.go b/logservice/logpuller/region_request_worker.go index 367ddc2561..614b67d7f0 100644 --- a/logservice/logpuller/region_request_worker.go +++ b/logservice/logpuller/region_request_worker.go @@ -25,172 +25,226 @@ import ( "github.com/pingcap/kvproto/pkg/kvrpcpb" "github.com/pingcap/log" cerror "github.com/pingcap/ticdc/pkg/errors" - "github.com/pingcap/ticdc/pkg/security" + "github.com/pingcap/ticdc/pkg/metrics" "github.com/pingcap/ticdc/pkg/util" "github.com/pingcap/ticdc/pkg/version" + "github.com/pingcap/ticdc/utils/notifyqueue" "go.uber.org/zap" "golang.org/x/sync/errgroup" grpcstatus "google.golang.org/grpc/status" ) +const storeReconnectBackoff = time.Second + // To generate a workerID in `newRegionRequestWorker`. var workerIDGen atomic.Uint64 -type regionFeedStates map[uint64]*regionFeedState +var ( + metricsResolvedTsCount = metrics.PullerEventCounter.WithLabelValues("resolved_ts") + metricBatchResolvedSize = metrics.BatchResolvedEventSize.WithLabelValues("event-store") +) -// regionRequestWorker is responsible for sending region requests to a specific TiKV store. -type regionRequestWorker struct { - workerID uint64 +type deregisterRequest struct { + subID SubscriptionID + filterLoop bool +} - client *subscriptionClient +type controlQueue struct { + mu sync.Mutex + queue *notifyqueue.Queue[deregisterRequest] +} - store *requestedStore +func newControlQueue() *controlQueue { + return &controlQueue{queue: notifyqueue.New[deregisterRequest]()} +} - // we must always get a region to request before create a grpc stream. - // only in this way we can avoid to try to connect to an offline store infinitely. - preFetchForConnecting *regionInfo +func (q *controlQueue) push(req deregisterRequest) { + q.mu.Lock() + defer q.mu.Unlock() + q.queue.Push(req) +} - // request cache with flow control - requestCache *requestCache +func (q *controlQueue) tryPop() (deregisterRequest, bool) { + q.mu.Lock() + defer q.mu.Unlock() + return q.queue.TryPop() +} - // all regions maintained by this worker. - requestedRegions struct { - sync.RWMutex +func (q *controlQueue) len() int { + q.mu.Lock() + defer q.mu.Unlock() + return q.queue.Len() +} - subscriptions map[SubscriptionID]regionFeedStates +func (q *controlQueue) drain() { + q.mu.Lock() + defer q.mu.Unlock() + for { + if _, ok := q.queue.TryPop(); !ok { + return + } } } +func (q *controlQueue) ready() <-chan struct{} { + return q.queue.Ready() +} + +// regionRequestWorker owns one TiKV event-feed stream and the requests sent +// through it, including reconnect cleanup and subscription deregistration. +type regionRequestWorker struct { + workerID uint64 + + upstream *upstreamHandle + eventSink *regionEventSink + failureHandler *regionFailureHandler + storeAddr string + + admission *regionAdmissionController + controlQueue *controlQueue + tracker *regionTracker +} + func newRegionRequestWorker( - ctx context.Context, - client *subscriptionClient, - credential *security.Credential, - g *errgroup.Group, - store *requestedStore, - requestCacheSize int, + upstream *upstreamHandle, + eventSink *regionEventSink, + failureHandler *regionFailureHandler, + storeAddr string, + currentWindow int, + maxWindowMultiplier int, + memoryQuota *memoryQuotaController, ) *regionRequestWorker { - worker := ®ionRequestWorker{ - workerID: workerIDGen.Add(1), - client: client, - store: store, - requestCache: newRequestCache(requestCacheSize), + workerID := workerIDGen.Add(1) + return ®ionRequestWorker{ + workerID: workerID, + upstream: upstream, + eventSink: eventSink, + failureHandler: failureHandler, + storeAddr: storeAddr, + admission: newRegionAdmissionController( + currentWindow, + maxWindowMultiplier, + memoryQuota, + upstream.pdClock, + ), + controlQueue: newControlQueue(), + tracker: newRegionTracker(), } - worker.requestedRegions.subscriptions = make(map[SubscriptionID]regionFeedStates) +} - waitForPreFetching := func() error { - if worker.preFetchForConnecting != nil { - log.Panic("preFetchForConnecting should be nil", - zap.Uint64("workerID", worker.workerID), - zap.String("addr", store.storeAddr)) +func (s *regionRequestWorker) Run(ctx context.Context) error { + handleStreamFailure := func(firstReq *regionReq, regionErr error) { + // Stream failure handle cases: + // - tracker: requests already sent to this stream. + // - firstReq: popped from admission for this stream, but not necessarily + // added to tracker yet if the stream fails before sendRegionRequest calls + // tracker.Add. + // - admission: requests owned by this worker but not sent yet. + for _, state := range s.tracker.Drain() { + state.markStopped(regionErr) + s.eventSink.Push( + SubscriptionID(state.requestID), + regionEvent{states: []*regionFeedState{state}}, + ) } - for { - req, err := worker.requestCache.pop(ctx) - if err != nil { - return err - } - if req.regionInfo.isStopped() { - worker.requestCache.markDone() - continue - } - worker.preFetchForConnecting = new(regionInfo) - *worker.preFetchForConnecting = req.regionInfo - return nil + // The failed stream no longer owns remote registrations. + s.controlQueue.drain() + if firstReq != nil && firstReq.abort() { + s.failureHandler.Report(newRegionErrorInfo(firstReq.regionInfo, regionErr)) + } + for _, task := range s.admission.drain() { + s.failureHandler.Report(newRegionErrorInfo(task.regionInfo, regionErr)) } } - g.Go(func() error { - for { - if err := waitForPreFetching(); err != nil { - return err - } - var regionErr error - if err := version.CheckStoreVersion(ctx, worker.client.pd); err != nil { - if errors.Cause(err) == context.Canceled { - return nil - } - log.Error("event feed check store version fails", - zap.Uint64("workerID", worker.workerID), - zap.String("addr", worker.store.storeAddr), - zap.Error(err)) - if cerror.Is(err, cerror.ErrGetAllStoresFailed) { - regionErr = &getStoreErr{} - } else { - regionErr = &storeStreamErr{} - } - } else { - if canceled := worker.run(ctx, credential); canceled { - return nil - } - regionErr = &storeStreamErr{} - } - for subID, m := range worker.clearRegionStates() { - for _, state := range m { - state.markStopped(regionErr) - regionEvent := regionEvent{ - states: []*regionFeedState{state}, - } - worker.client.pushRegionEventToDS(subID, regionEvent) - } - } - // The store may fail forever, so we need try to re-schedule all pending regions. - for _, region := range worker.clearPendingRegions() { - if region.isStopped() { - // It means it's a special task for stopping the table. - continue - } - client.onRegionFail(newRegionErrorInfo(region, regionErr)) - } - if err := util.Hang(ctx, time.Second); err != nil { - return err - } + for { + // Do not connect an idle worker to an unavailable store indefinitely. + firstReq, err := s.waitForRegionRequest(ctx) + if err != nil { + return err } - }) - return worker + regionErr := s.runStream(ctx, firstReq) + if ctx.Err() != nil { + firstReq.abort() + return ctx.Err() + } + // Treat an unexpected clean stream exit as a recoverable store-stream failure. + if regionErr == nil { + regionErr = &storeStreamErr{} + } + handleStreamFailure(firstReq, regionErr) + if err := util.Hang(ctx, storeReconnectBackoff); err != nil { + return err + } + } } -func (s *regionRequestWorker) run(ctx context.Context, credential *security.Credential) (canceled bool) { - isCanceled := func() bool { - select { - case <-ctx.Done(): - return true - default: - return false - } +func (s *regionRequestWorker) waitForRegionRequest(ctx context.Context) (*regionReq, error) { + // Without a stream there are no remote registrations to deregister. + s.controlQueue.drain() + req, err := s.admission.pop(ctx, nil) + if err != nil { + return nil, err } + // Drop controls that raced with selecting the first request. Any later + // controls will be handled by the stream send loop. + s.controlQueue.drain() + return req, nil +} - log.Info("region request worker going to create grpc stream", +func (s *regionRequestWorker) checkStoreVersion(ctx context.Context) error { + err := version.CheckStoreVersion(ctx, s.upstream.pd) + if err == nil { + return nil + } + if ctx.Err() != nil { + return ctx.Err() + } + log.Error("event feed check store version fails", zap.Uint64("workerID", s.workerID), - zap.String("addr", s.store.storeAddr)) + zap.String("addr", s.storeAddr), + zap.Error(err)) + if cerror.Is(err, cerror.ErrGetAllStoresFailed) { + return &getStoreErr{} + } + return &storeStreamErr{} +} +func (s *regionRequestWorker) runStream(ctx context.Context, firstReq *regionReq) (err error) { + if err := s.checkStoreVersion(ctx); err != nil { + return err + } + + log.Info("region request worker going to create grpc stream", + zap.Uint64("workerID", s.workerID), + zap.String("addr", s.storeAddr)) defer func() { log.Info("region request worker exits", zap.Uint64("workerID", s.workerID), - zap.String("addr", s.store.storeAddr), - zap.Bool("canceled", canceled)) + zap.String("addr", s.storeAddr), + zap.Error(err)) }() - g, gctx := errgroup.WithContext(ctx) - conn, err := Connect(gctx, credential, s.store.storeAddr) + conn, err := Connect(ctx, s.upstream.credential, s.storeAddr) if err != nil { log.Warn("region request worker create grpc stream failed", zap.Uint64("workerID", s.workerID), - zap.String("addr", s.store.storeAddr), + zap.String("addr", s.storeAddr), zap.Error(err)) - // Close the connection if it was partially created to prevent goroutine leaks if conn != nil && conn.Conn != nil { _ = conn.Conn.Close() } - return isCanceled() + if ctx.Err() != nil { + return ctx.Err() + } + return &storeStreamErr{} } - defer func() { - _ = conn.Conn.Close() - }() + defer func() { _ = conn.Conn.Close() }() - g.Go(func() error { - return s.receiveAndDispatchChangeEvents(conn) - }) - g.Go(func() error { return s.processRegionSendTask(gctx, conn) }) + g, gctx := errgroup.WithContext(ctx) + g.Go(func() error { return s.receiveAndDispatchChangeEvents(conn) }) + g.Go(func() error { return s.processRegionSendTask(gctx, conn, firstReq) }) failpoint.Inject("InjectForceReconnect", func() { timer := time.After(10 * time.Second) @@ -202,8 +256,14 @@ func (s *regionRequestWorker) run(ctx context.Context, credential *security.Cred }) }) - _ = g.Wait() - return isCanceled() + err = g.Wait() + if err != nil { + if ctx.Err() != nil { + return ctx.Err() + } + return &storeStreamErr{} + } + return nil } func normalizeStreamError(err error) error { @@ -213,14 +273,13 @@ func normalizeStreamError(err error) error { return errors.Trace(err) } -// receiveAndDispatchChangeEventsToProcessor receives events from the grpc stream and dispatches them to ds. func (s *regionRequestWorker) receiveAndDispatchChangeEvents(conn *ConnAndClient) error { for { changeEvent, err := conn.Client.Recv() if err != nil { log.Info("region request worker receive from grpc stream failed", zap.Uint64("workerID", s.workerID), - zap.String("addr", s.store.storeAddr), + zap.String("addr", s.storeAddr), zap.String("code", grpcstatus.Code(err).String()), zap.Error(err)) return normalizeStreamError(err) @@ -238,11 +297,9 @@ func (s *regionRequestWorker) dispatchRegionChangeEvents(events []*cdcpb.Event) for _, event := range events { regionID := event.RegionId subscriptionID := SubscriptionID(event.RequestId) - state := s.getRegionState(subscriptionID, regionID) + state := s.tracker.Get(subscriptionID, regionID) if state != nil { - regionEvent := regionEvent{ - states: []*regionFeedState{state}, - } + regionEvent := regionEvent{states: []*regionFeedState{state}} switch eventData := event.Event.(type) { case *cdcpb.Event_Entries_: if eventData == nil { @@ -266,12 +323,11 @@ func (s *regionRequestWorker) dispatchRegionChangeEvents(events []*cdcpb.Event) case *cdcpb.Event_ResolvedTs: regionEvent.resolvedTs = eventData.ResolvedTs case *cdcpb.Event_LongTxn_: - // ignore continue default: log.Panic("unknown event type", zap.Any("event", event)) } - s.client.pushRegionEventToDS(subscriptionID, regionEvent) + s.eventSink.Push(subscriptionID, regionEvent) } else { switch event.Event.(type) { case *cdcpb.Event_Error: @@ -293,7 +349,7 @@ func (s *regionRequestWorker) dispatchRegionChangeEvents(events []*cdcpb.Event) func (s *regionRequestWorker) dispatchResolvedTsEvent(resolvedTsEvent *cdcpb.ResolvedTs) { subscriptionID := SubscriptionID(resolvedTsEvent.RequestId) metricsResolvedTsCount.Add(float64(len(resolvedTsEvent.Regions))) - s.client.metrics.batchResolvedSize.Observe(float64(len(resolvedTsEvent.Regions))) + metricBatchResolvedSize.Observe(float64(len(resolvedTsEvent.Regions))) // TODO: resolvedTsEvent.Ts be 0 is impossible, we need find the root cause. if resolvedTsEvent.Ts == 0 { log.Warn("region request worker receives a resolved ts event with zero value, ignore it", @@ -302,34 +358,29 @@ func (s *regionRequestWorker) dispatchResolvedTsEvent(resolvedTsEvent *cdcpb.Res zap.Any("regionIDs", resolvedTsEvent.Regions)) return } + + const resolvedTsStateBatchSize = 1024 // Avoid allocating a huge states slice when resolvedTsEvent.Regions is large. // Push resolved-ts events in batches to reduce peak memory usage and improve GC behavior. - const resolvedTsStateBatchSize = 1024 - capHint := len(resolvedTsEvent.Regions) - if capHint > resolvedTsStateBatchSize { - capHint = resolvedTsStateBatchSize - } + capHint := min(len(resolvedTsEvent.Regions), resolvedTsStateBatchSize) resolvedStates := make([]*regionFeedState, 0, capHint) flush := func() { if len(resolvedStates) == 0 { return } - s.client.pushRegionEventToDS(subscriptionID, regionEvent{ + s.eventSink.Push(subscriptionID, regionEvent{ resolvedTs: resolvedTsEvent.Ts, states: resolvedStates, }) resolvedStates = nil } for i, regionID := range resolvedTsEvent.Regions { - if state := s.getRegionState(subscriptionID, regionID); state != nil { + if state := s.tracker.Get(subscriptionID, regionID); state != nil { resolvedStates = append(resolvedStates, state) if len(resolvedStates) >= resolvedTsStateBatchSize { flush() if i+1 < len(resolvedTsEvent.Regions) { - capHint = len(resolvedTsEvent.Regions) - (i + 1) - if capHint > resolvedTsStateBatchSize { - capHint = resolvedTsStateBatchSize - } + capHint = min(len(resolvedTsEvent.Regions)-(i+1), resolvedTsStateBatchSize) resolvedStates = make([]*regionFeedState, 0, capHint) } } @@ -344,102 +395,129 @@ func (s *regionRequestWorker) dispatchResolvedTsEvent(resolvedTsEvent *cdcpb.Res flush() } -// processRegionSendTask receives region requests from the channel and sends them to the remote store. -func (s *regionRequestWorker) processRegionSendTask( - ctx context.Context, +func (s *regionRequestWorker) sendChangeDataRequest( conn *ConnAndClient, + req *cdcpb.ChangeDataRequest, ) error { - doSend := func(req *cdcpb.ChangeDataRequest) error { - if err := conn.Client.Send(req); err != nil { - log.Warn("region request worker send request to grpc stream failed", - zap.Uint64("workerID", s.workerID), - zap.Uint64("subscriptionID", req.RequestId), - zap.Uint64("regionID", req.RegionId), - zap.String("addr", s.store.storeAddr), - zap.Error(err)) - return normalizeStreamError(err) - } - // TODO: add a metric? + if err := conn.Client.Send(req); err != nil { + log.Warn("region request worker send request to grpc stream failed", + zap.Uint64("workerID", s.workerID), + zap.Uint64("subscriptionID", req.RequestId), + zap.Uint64("regionID", req.RegionId), + zap.String("addr", s.storeAddr), + zap.Error(err)) + return normalizeStreamError(err) + } + return nil +} + +func (s *regionRequestWorker) sendDeregisterRequest( + conn *ConnAndClient, + req deregisterRequest, +) error { + changeDataReq := &cdcpb.ChangeDataRequest{ + Header: &cdcpb.Header{ClusterId: s.upstream.clusterID, TicdcVersion: version.ReleaseSemver()}, + RequestId: uint64(req.subID), + Request: &cdcpb.ChangeDataRequest_Deregister_{ + Deregister: &cdcpb.ChangeDataRequest_Deregister{}, + }, + FilterLoop: req.filterLoop, + } + if err := s.sendChangeDataRequest(conn, changeDataReq); err != nil { + return err + } + for _, state := range s.tracker.TakeSubscription(req.subID) { + state.markStopped(&requestCancelledErr{}) + s.eventSink.Push(req.subID, regionEvent{states: []*regionFeedState{state}}) + } + return nil +} + +func (s *regionRequestWorker) sendRegionRequest(conn *ConnAndClient, req *regionReq) error { + if !req.isActive() { + return nil + } + region := req.regionInfo + subID := region.subscribedSpan.subID + log.Debug("region request worker sends region request", + zap.Uint64("workerID", s.workerID), + zap.Uint64("subscriptionID", uint64(subID)), + zap.Uint64("regionID", region.verID.GetID()), + zap.String("addr", s.storeAddr), + zap.Bool("bdrMode", region.filterLoop)) + + if region.subscribedSpan.stopped.Load() { + req.abort() + s.failureHandler.Report(newRegionErrorInfo(region, &requestCancelledErr{})) return nil } - // Handle pre-fetched region first - region := *s.preFetchForConnecting - s.preFetchForConnecting = nil - regionReq := newRegionReq(region) - var err error - for { - region := regionReq.regionInfo - subID := region.subscribedSpan.subID - log.Debug("region request worker gets a singleRegionInfo", + // Publish the state before Send so a fast response observes its owner and + // admission lease. + state := newRegionFeedState(region, uint64(subID), s, req, func(state *regionFeedState) { + s.failureHandler.resetRegionRecovery(state.region) + }) + if !s.tracker.Add(subID, region.verID.GetID(), state) { + // RangeLock normally prevents duplicate active regions. Keep the existing + // owner, including its range-lock ownership, if that invariant is ever + // violated. Only the duplicate request's flow-control slot is released. + state.abortScanIfNeeded() + state.matcher.clear() + log.Warn("duplicate active region request ignored", zap.Uint64("workerID", s.workerID), zap.Uint64("subscriptionID", uint64(subID)), - zap.Uint64("regionID", region.verID.GetID()), - zap.String("addr", s.store.storeAddr), - zap.Bool("bdrMode", region.filterLoop)) - - // It means it's a special task for stopping the table. - if region.isStopped() { - req := &cdcpb.ChangeDataRequest{ - Header: &cdcpb.Header{ClusterId: s.client.clusterID, TicdcVersion: version.ReleaseSemver()}, - RequestId: uint64(subID), - Request: &cdcpb.ChangeDataRequest_Deregister_{ - Deregister: &cdcpb.ChangeDataRequest_Deregister{}, - }, - FilterLoop: region.filterLoop, - } - s.requestCache.markDone() - if err := doSend(req); err != nil { + zap.Uint64("regionID", region.verID.GetID())) + return nil + } + if err := s.sendChangeDataRequest(conn, createRegionRequest(s.upstream.clusterID, region)); err != nil { + // Transport failures are always recoverable at the region level. Preserve + // the stream error as the function result, but classify the region for + // rescheduling instead of exposing an arbitrary gRPC error downstream. + state.markStopped(&storeStreamErr{}) + return err + } + return nil +} + +func (s *regionRequestWorker) processRegionSendTask( + ctx context.Context, + conn *ConnAndClient, + firstReq *regionReq, +) error { + regionReq := firstReq + for { + // Send the current region request before handling anything newly queued. + if regionReq != nil { + if err := s.sendRegionRequest(conn, regionReq); err != nil { return err } - for _, state := range s.takeRegionStates(subID) { - state.markStopped(&requestCancelledErr{}) - regionEvent := regionEvent{ - states: []*regionFeedState{state}, - } - s.client.pushRegionEventToDS(subID, regionEvent) + } + // Flush pending deregisters before admitting the next region request. + // Admission may still contain stale tasks from a stopped subscription, but + // sendRegionRequest re-checks subscription liveness before tracker.Add/Send, + // so those tasks are dropped locally instead of recreating remote registrations. + for { + req, ok := s.controlQueue.tryPop() + if !ok { + break } - } else if region.subscribedSpan.stopped.Load() { - // It can be skipped directly because there must be no pending states from - // the stopped subscribedTable, or the special singleRegionInfo for stopping - // the table will be handled later. - s.client.onRegionFail(newRegionErrorInfo(region, &storeStreamErr{})) - s.requestCache.markDone() - } else { - state := newRegionFeedState(region, uint64(subID), s) - state.start() - s.addRegionState(subID, region.verID.GetID(), state) - // Mark the request as sent before sending it. - // Otherwise there is a race with the receiver goroutine: - // 1. addRegionState makes the region visible to error handling. - // 2. doSend sends the request. - // 3. the receiver goroutine may receive a region error immediately. - // 4. markStopped runs before markSent, so requestCache.markStopped cannot - // find the request in sentRequests. - // 5. the sender goroutine then calls markSent and leaves a stale sent - // request behind, even though the region has already been - // unlocked/rescheduled. - // - // Tracking the request before Send keeps requestedRegions and - // sentRequests visible in the same order and avoids leaving stale - // requests in cleanup. - s.requestCache.markSent(regionReq) - if err := doSend(s.createRegionRequest(region)); err != nil { - state.markStopped(err) + if err := s.sendDeregisterRequest(conn, req); err != nil { return err } } - // Try to get from cache - regionReq, err = s.requestCache.pop(ctx) + // Block for the next request, but wake early when deregisters arrive. + // regionReq above is already consumed and will be replaced by the next pop. + var err error + regionReq, err = s.admission.pop(ctx, s.controlQueue.ready()) if err != nil { return err } } } -func (s *regionRequestWorker) createRegionRequest(region regionInfo) *cdcpb.ChangeDataRequest { +func createRegionRequest(clusterID uint64, region regionInfo) *cdcpb.ChangeDataRequest { return &cdcpb.ChangeDataRequest{ - Header: &cdcpb.Header{ClusterId: s.client.clusterID, TicdcVersion: version.ReleaseSemver()}, + Header: &cdcpb.Header{ClusterId: clusterID, TicdcVersion: version.ReleaseSemver()}, RegionId: region.verID.GetID(), RequestId: uint64(region.subscribedSpan.subID), RegionEpoch: region.rpcCtx.Meta.RegionEpoch, @@ -451,79 +529,3 @@ func (s *regionRequestWorker) createRegionRequest(region regionInfo) *cdcpb.Chan ScanPriority: normalizeScanPriority(region.scanPriority), } } - -func (s *regionRequestWorker) addRegionState(subscriptionID SubscriptionID, regionID uint64, state *regionFeedState) { - s.requestedRegions.Lock() - defer s.requestedRegions.Unlock() - states := s.requestedRegions.subscriptions[subscriptionID] - if states == nil { - states = make(regionFeedStates) - s.requestedRegions.subscriptions[subscriptionID] = states - } - - states[regionID] = state -} - -func (s *regionRequestWorker) getRegionState(subscriptionID SubscriptionID, regionID uint64) *regionFeedState { - s.requestedRegions.RLock() - defer s.requestedRegions.RUnlock() - if states, ok := s.requestedRegions.subscriptions[subscriptionID]; ok { - return states[regionID] - } - return nil -} - -func (s *regionRequestWorker) takeRegionState(subscriptionID SubscriptionID, regionID uint64) *regionFeedState { - s.requestedRegions.Lock() - defer s.requestedRegions.Unlock() - if statesMap, ok := s.requestedRegions.subscriptions[subscriptionID]; ok { - state := statesMap[regionID] - delete(statesMap, regionID) - if len(statesMap) == 0 { - delete(s.requestedRegions.subscriptions, subscriptionID) - } - return state - } - return nil -} - -func (s *regionRequestWorker) takeRegionStates(subscriptionID SubscriptionID) regionFeedStates { - s.requestedRegions.Lock() - defer s.requestedRegions.Unlock() - states := s.requestedRegions.subscriptions[subscriptionID] - delete(s.requestedRegions.subscriptions, subscriptionID) - return states -} - -func (s *regionRequestWorker) clearRegionStates() map[SubscriptionID]regionFeedStates { - s.requestedRegions.Lock() - defer s.requestedRegions.Unlock() - subscriptions := s.requestedRegions.subscriptions - s.requestedRegions.subscriptions = make(map[SubscriptionID]regionFeedStates) - return subscriptions -} - -// add adds a region request to the worker's cache -// It blocks if the cache is full until there's space or ctx is cancelled -func (s *regionRequestWorker) add(ctx context.Context, region regionInfo, force bool) (bool, error) { - return s.requestCache.add(ctx, region, force) -} - -func (s *regionRequestWorker) clearPendingRegions() []regionInfo { - var regions []regionInfo - - // Clear pre-fetched region - if s.preFetchForConnecting != nil { - region := *s.preFetchForConnecting - s.preFetchForConnecting = nil - regions = append(regions, region) - // The pre-fetched region was popped from pendingQueue but hasn't been marked as sent or done yet. - // Release its pendingCount slot to avoid leaking flow control credits on worker failures. - s.requestCache.markDone() - } - - // Clear all regions from cache - cacheRegions := s.requestCache.clear() - regions = append(regions, cacheRegions...) - return regions -} diff --git a/logservice/logpuller/region_request_worker_test.go b/logservice/logpuller/region_request_worker_test.go index 26b2f923ae..4f429925ee 100644 --- a/logservice/logpuller/region_request_worker_test.go +++ b/logservice/logpuller/region_request_worker_test.go @@ -17,14 +17,19 @@ import ( "context" "io" "testing" + "time" "github.com/pingcap/errors" "github.com/pingcap/kvproto/pkg/cdcpb" "github.com/pingcap/kvproto/pkg/metapb" + "github.com/pingcap/ticdc/heartbeatpb" "github.com/pingcap/ticdc/logservice/logpuller/regionlock" + "github.com/pingcap/ticdc/pkg/security" "github.com/pingcap/ticdc/utils/dynstream" - "github.com/prometheus/client_golang/prometheus" + "github.com/pingcap/tidb/pkg/store/mockstore/mockcopr" "github.com/stretchr/testify/require" + "github.com/tikv/client-go/v2/oracle" + "github.com/tikv/client-go/v2/testutils" "github.com/tikv/client-go/v2/tikv" "google.golang.org/grpc" "google.golang.org/grpc/codes" @@ -35,9 +40,15 @@ import ( type mockEventFeedV2Client struct { sendErr error recvErr error + sendCh chan *cdcpb.ChangeDataRequest } -func (m *mockEventFeedV2Client) Send(*cdcpb.ChangeDataRequest) error { return m.sendErr } +func (m *mockEventFeedV2Client) Send(req *cdcpb.ChangeDataRequest) error { + if m.sendCh != nil { + m.sendCh <- req + } + return m.sendErr +} func (m *mockEventFeedV2Client) Recv() (*cdcpb.ChangeDataEvent, error) { return nil, m.recvErr } func (m *mockEventFeedV2Client) Header() (metadata.MD, error) { return metadata.MD{}, nil } func (m *mockEventFeedV2Client) Trailer() metadata.MD { return metadata.MD{} } @@ -58,10 +69,6 @@ func prepareRegionForSendTest(region regionInfo) regionInfo { } func TestCreateRegionRequestScanPriority(t *testing.T) { - worker := ®ionRequestWorker{ - client: &subscriptionClient{clusterID: 1}, - } - for _, tc := range []struct { name string priority cdcpb.ScanPriority @@ -87,55 +94,151 @@ func TestCreateRegionRequestScanPriority(t *testing.T) { region := prepareRegionForSendTest(createTestRegionInfo(1, 1)) region.scanPriority = tc.priority - req := worker.createRegionRequest(region) + req := createRegionRequest(1, region) require.Equal(t, tc.expected, req.GetScanPriority()) }) } } -func TestRegionStatesOperation(t *testing.T) { - worker := ®ionRequestWorker{} - worker.requestedRegions.subscriptions = make(map[SubscriptionID]regionFeedStates) +func admitRegionRequest( + t *testing.T, + controller *regionAdmissionController, + region regionInfo, +) *regionReq { + t.Helper() + currentTs := oracle.GoTimeToTS(time.Now()) + submitRegionForAdmission(t, controller, region, currentTs) + req, err := controller.pop(t.Context(), nil) + require.NoError(t, err) + return req +} - require.Nil(t, worker.getRegionState(1, 2)) - require.Nil(t, worker.takeRegionState(1, 2)) +type pushedRegionEvent struct { + subscriptionID SubscriptionID + event regionEvent +} + +type recordingRegionEventDynamicStream struct { + events chan pushedRegionEvent +} + +func (m *recordingRegionEventDynamicStream) Start() {} - worker.addRegionState(1, 2, ®ionFeedState{}) - require.NotNil(t, worker.getRegionState(1, 2)) - require.NotNil(t, worker.takeRegionState(1, 2)) - require.Nil(t, worker.getRegionState(1, 2)) - require.Equal(t, 0, len(worker.requestedRegions.subscriptions)) +func (m *recordingRegionEventDynamicStream) Close() {} - worker.addRegionState(1, 2, ®ionFeedState{}) - require.NotNil(t, worker.getRegionState(1, 2)) - require.NotNil(t, worker.takeRegionState(1, 2)) - require.Nil(t, worker.getRegionState(1, 2)) - require.Equal(t, 0, len(worker.requestedRegions.subscriptions)) +func (m *recordingRegionEventDynamicStream) Push(path SubscriptionID, event regionEvent) { + m.events <- pushedRegionEvent{subscriptionID: path, event: event} } -func TestClearPendingRegionsReleaseSlotForPreFetchedRegion(t *testing.T) { - worker := ®ionRequestWorker{ - requestCache: newRequestCache(10), +func (m *recordingRegionEventDynamicStream) Wake(SubscriptionID) {} + +func (m *recordingRegionEventDynamicStream) Feedback() <-chan dynstream.Feedback[int, SubscriptionID, *subscribedSpan] { + return nil +} + +func (m *recordingRegionEventDynamicStream) AddPath(SubscriptionID, *subscribedSpan, ...dynstream.AreaSettings) error { + return nil +} + +func (m *recordingRegionEventDynamicStream) RemovePath(SubscriptionID) error { + return nil +} + +func (m *recordingRegionEventDynamicStream) Release(SubscriptionID) {} + +func (m *recordingRegionEventDynamicStream) SetAreaSettings(int, dynstream.AreaSettings) {} + +func (m *recordingRegionEventDynamicStream) GetMetrics() dynstream.Metrics[int, SubscriptionID] { + return dynstream.Metrics[int, SubscriptionID]{} +} + +func createFailureRecoveryTestRegion(t *testing.T, subID SubscriptionID, regionID uint64) regionInfo { + t.Helper() + + fullSpan := heartbeatpb.TableSpan{ + TableID: 1, + StartKey: []byte("a"), + EndKey: []byte("z"), + } + subSpan := &subscribedSpan{ + subID: subID, + startTs: 100, + span: fullSpan, + rangeLock: regionlock.NewRangeLock(1, fullSpan.StartKey, fullSpan.EndKey, 100), } + regionSpan := heartbeatpb.TableSpan{ + TableID: 1, + StartKey: []byte("a"), + EndKey: []byte("m"), + } + locked := subSpan.rangeLock.LockRange(context.Background(), regionSpan.StartKey, regionSpan.EndKey, regionID, 1) + require.Equal(t, regionlock.LockRangeStatusSuccess, locked.Status) + other := subSpan.rangeLock.LockRange(context.Background(), []byte("m"), []byte("z"), regionID+1000, 1) + require.Equal(t, regionlock.LockRangeStatusSuccess, other.Status) - ctx := context.Background() - region := createTestRegionInfo(1, 1) + region := newRegionInfo(tikv.NewRegionVerID(regionID, 1, 1), regionSpan, nil, subSpan, false) + region.lockedRangeState = locked.LockedRangeState + region.lockedRangeState.ResolvedTs.Store(100) + return region +} - ok, err := worker.requestCache.add(ctx, region, false) - require.NoError(t, err) - require.True(t, ok) +func newFailureRecoveryTestPDClient(t *testing.T) *mockPDClient { + t.Helper() + + _, _, pdClient, _ := testutils.NewMockTiKV("", mockcopr.NewCoprRPCHandler()) + return &mockPDClient{Client: pdClient, versionGen: defaultVersionGen} +} + +func snapshotErrCacheRegionIDs(handler *regionFailureHandler) []uint64 { + handler.cache.Lock() + defer handler.cache.Unlock() + + regionIDs := make([]uint64, 0, len(handler.cache.cache)) + for _, errInfo := range handler.cache.cache { + regionIDs = append(regionIDs, errInfo.verID.GetID()) + } + return regionIDs +} + +func errCacheLen(handler *regionFailureHandler) int { + handler.cache.Lock() + defer handler.cache.Unlock() + return len(handler.cache.cache) +} - req, err := worker.requestCache.pop(ctx) +func TestRegionRequestWorkerIgnoresDuplicateActiveRegion(t *testing.T) { + admission := newTestRegionAdmissionController(10, 1) + worker := ®ionRequestWorker{ + admission: admission, + storeAddr: "store-1", + upstream: &upstreamHandle{}, + tracker: newRegionTracker(), + } + region := prepareRegionForSendTest(createTestRegionInfo(1, 1)) + + req1 := admitRegionRequest(t, admission, region) + state1 := newRegionFeedState(region, uint64(region.subscribedSpan.subID), worker, req1, nil) + require.True(t, worker.tracker.Add(region.subscribedSpan.subID, region.verID.GetID(), state1)) + + req2 := admitRegionRequest(t, admission, region) + sendCh := make(chan *cdcpb.ChangeDataRequest, 1) + err := worker.sendRegionRequest(&ConnAndClient{ + Client: &mockEventFeedV2Client{sendCh: sendCh}, + Conn: &grpc.ClientConn{}, + }, req2) require.NoError(t, err) - require.Equal(t, 1, worker.requestCache.getPendingCount()) - worker.preFetchForConnecting = new(regionInfo) - *worker.preFetchForConnecting = req.regionInfo + require.Equal(t, 1, admission.stats().inflight) + require.Same(t, state1, worker.tracker.Get(region.subscribedSpan.subID, region.verID.GetID())) + require.False(t, state1.isStale()) + select { + case <-sendCh: + t.Fatal("duplicate region request must not be sent") + default: + } - regions := worker.clearPendingRegions() - require.Len(t, regions, 1) - require.Nil(t, worker.preFetchForConnecting) - require.Equal(t, 0, worker.requestCache.getPendingCount()) + state1.abortScanIfNeeded() + state1.matcher.clear() } type pushedResolvedEvent struct { @@ -150,6 +253,15 @@ type mockRegionEventDynamicStream struct { pushed []pushedResolvedEvent } +type countingRegionEventDynamicStream struct { + mockRegionEventDynamicStream + pushCount int +} + +func (m *countingRegionEventDynamicStream) Push(_ SubscriptionID, _ regionEvent) { + m.pushCount++ +} + func (m *mockRegionEventDynamicStream) Start() {} func (m *mockRegionEventDynamicStream) Close() {} @@ -189,23 +301,16 @@ func (m *mockRegionEventDynamicStream) GetMetrics() dynstream.Metrics[int, Subsc func newDispatchResolvedTsTestWorker(regionCount int) (*regionRequestWorker, *mockRegionEventDynamicStream, *cdcpb.ResolvedTs) { ds := &mockRegionEventDynamicStream{} worker := ®ionRequestWorker{ - client: &subscriptionClient{ - metrics: sharedClientMetrics{ - batchResolvedSize: prometheus.ObserverFunc(func(float64) {}), - }, - eventSink: newTestRegionEventSink(ds), - }, - } - worker.requestedRegions.subscriptions = map[SubscriptionID]regionFeedStates{ - 1: make(regionFeedStates, regionCount), + eventSink: ®ionEventSink{ds: ds}, + tracker: newRegionTracker(), } regions := make([]uint64, regionCount) for i := 0; i < regionCount; i++ { regionID := uint64(i + 1) regions[i] = regionID - worker.requestedRegions.subscriptions[1][regionID] = ®ionFeedState{ + worker.tracker.Add(1, regionID, ®ionFeedState{ requestID: 1, - } + }) } return worker, ds, &cdcpb.ResolvedTs{ @@ -224,14 +329,14 @@ func dispatchResolvedTsEventLegacyForBenchmark(s *regionRequestWorker, resolvedT return } states := resolvedStates - s.client.pushRegionEventToDS(subscriptionID, regionEvent{ + s.eventSink.Push(subscriptionID, regionEvent{ resolvedTs: resolvedTsEvent.Ts, states: states, }) resolvedStates = make([]*regionFeedState, 0, resolvedTsStateBatchSize) } for _, regionID := range resolvedTsEvent.Regions { - if state := s.getRegionState(subscriptionID, regionID); state != nil { + if state := s.tracker.Get(subscriptionID, regionID); state != nil { resolvedStates = append(resolvedStates, state) if len(resolvedStates) >= resolvedTsStateBatchSize { flush() @@ -242,7 +347,9 @@ func dispatchResolvedTsEventLegacyForBenchmark(s *regionRequestWorker, resolvedT } func benchmarkDispatchResolvedTsEvent(b *testing.B, regionCount int, useLegacy bool) { - worker, ds, event := newDispatchResolvedTsTestWorker(regionCount) + worker, _, event := newDispatchResolvedTsTestWorker(regionCount) + ds := &countingRegionEventDynamicStream{} + worker.eventSink.ds = ds b.ReportAllocs() b.ResetTimer() for i := 0; i < b.N; i++ { @@ -258,6 +365,41 @@ func benchmarkDispatchResolvedTsEvent(b *testing.B, regionCount int, useLegacy b } } +func TestWaitForRegionRequestDrainsIdleControlQueue(t *testing.T) { + admission := newTestRegionAdmissionController(1, 1) + worker := ®ionRequestWorker{ + admission: admission, + controlQueue: newControlQueue(), + } + + type waitResult struct { + req *regionReq + err error + } + resultCh := make(chan waitResult, 1) + go func() { + req, err := worker.waitForRegionRequest(t.Context()) + resultCh <- waitResult{req: req, err: err} + }() + + for subID := SubscriptionID(1); subID <= 100; subID++ { + worker.controlQueue.push(deregisterRequest{subID: subID}) + } + + region := prepareRegionForAdmission(createTestRegionInfo(1, 1), 1) + submitRegionForAdmission(t, admission, region, 1) + + select { + case result := <-resultCh: + require.NoError(t, result.err) + require.NotNil(t, result.req) + require.Equal(t, uint64(1), result.req.regionInfo.verID.GetID()) + case <-time.After(time.Second): + t.Fatal("worker did not receive the first region request") + } + require.Zero(t, worker.controlQueue.len()) +} + func TestDispatchResolvedTsEventSingleRegion(t *testing.T) { worker, ds, event := newDispatchResolvedTsTestWorker(1) worker.dispatchResolvedTsEvent(event) @@ -305,59 +447,123 @@ func BenchmarkDispatchResolvedTsEventSmallBatchCurrent(b *testing.B) { benchmarkDispatchResolvedTsEvent(b, 16, false) } -func TestClearPendingRegionsDoesNotReturnStoppedSentRegion(t *testing.T) { +func TestStoppedStateRemovesSentRequest(t *testing.T) { + admission := newTestRegionAdmissionController(10, 1) worker := ®ionRequestWorker{ - requestCache: newRequestCache(10), + admission: admission, + tracker: newRegionTracker(), } - worker.requestedRegions.subscriptions = make(map[SubscriptionID]regionFeedStates) + region := prepareRegionForSendTest(createTestRegionInfo(1, 1)) + req := admitRegionRequest(t, admission, region) - ctx := context.Background() - region := createTestRegionInfo(1, 1) + state := newRegionFeedState(req.regionInfo, uint64(req.regionInfo.subscribedSpan.subID), worker, req, nil) + require.True(t, worker.tracker.Add(req.regionInfo.subscribedSpan.subID, req.regionInfo.verID.GetID(), state)) + state.markStopped(errors.New("send request to store error")) + worker.tracker.RemoveIf(req.regionInfo.subscribedSpan.subID, req.regionInfo.verID.GetID(), state) - ok, err := worker.requestCache.add(ctx, region, false) - require.NoError(t, err) - require.True(t, ok) + require.Equal(t, 0, admission.stats().inflight) +} - req, err := worker.requestCache.pop(ctx) - require.NoError(t, err) +func TestRunStreamFailurePushesTrackedRegionToEventSink(t *testing.T) { + pdClient := newFailureRecoveryTestPDClient(t) + defer pdClient.Close() - state := newRegionFeedState(req.regionInfo, uint64(req.regionInfo.subscribedSpan.subID), worker) - state.start() - worker.addRegionState(req.regionInfo.subscribedSpan.subID, req.regionInfo.verID.GetID(), state) + ds := &recordingRegionEventDynamicStream{events: make(chan pushedRegionEvent, 4)} + handler := newRegionFailureHandler(nil, func(*subscribedSpan) {}, nil, nil) + worker := ®ionRequestWorker{ + upstream: &upstreamHandle{pd: pdClient, credential: &security.Credential{}}, + eventSink: ®ionEventSink{ds: ds}, + failureHandler: handler, + admission: newTestRegionAdmissionController(10, 1), + controlQueue: newControlQueue(), + tracker: newRegionTracker(), + storeAddr: "127.0.0.1:1", + } - // Simulate the race we are fixing in processRegionSendTask: - // once a request is visible in sentRequests, a fast region error may mark the - // region stopped before worker cleanup runs. In that case, markStopped should - // remove the sent request immediately, so clearPendingRegions must not return - // the stale region again during worker shutdown. - worker.requestCache.markSent(req) - state.markStopped(errors.New("send request to store error")) - worker.takeRegionState(req.regionInfo.subscribedSpan.subID, req.regionInfo.verID.GetID()) + sentRegion := createFailureRecoveryTestRegion(t, 1, 1) + sentReq := admitRegionRequest(t, worker.admission, sentRegion) + sentState := newRegionFeedState(sentRegion, uint64(sentRegion.subscribedSpan.subID), worker, sentReq, nil) + require.True(t, worker.tracker.Add(sentRegion.subscribedSpan.subID, sentRegion.verID.GetID(), sentState)) + + firstRegion := createFailureRecoveryTestRegion(t, 2, 2) + submitRegionForAdmission(t, worker.admission, firstRegion, 100) + + ctx, cancel := context.WithCancel(context.Background()) + runErrCh := make(chan error, 1) + go func() { + runErrCh <- worker.Run(ctx) + }() + + var pushed pushedRegionEvent + select { + case pushed = <-ds.events: + case <-time.After(5 * time.Second): + t.Fatal("worker did not push tracked region after stream failure") + } - require.Equal(t, 0, worker.requestCache.getPendingCount()) - require.Empty(t, worker.clearPendingRegions()) + require.Equal(t, SubscriptionID(1), pushed.subscriptionID) + require.Len(t, pushed.event.states, 1) + require.Same(t, sentState, pushed.event.states[0]) + require.Equal(t, 0, worker.admission.stats().inflight) + + var streamErr *storeStreamErr + require.ErrorAs(t, sentState.takeError(), &streamErr) + + cancel() + require.ErrorIs(t, <-runErrCh, context.Canceled) +} + +func TestRunStreamFailureReportsPendingRegionsToFailureHandler(t *testing.T) { + pdClient := newFailureRecoveryTestPDClient(t) + defer pdClient.Close() + + handler := newRegionFailureHandler(nil, func(*subscribedSpan) {}, nil, nil) + worker := ®ionRequestWorker{ + upstream: &upstreamHandle{pd: pdClient, credential: &security.Credential{}}, + eventSink: ®ionEventSink{ds: &mockDynamicStream{}}, + failureHandler: handler, + admission: newTestRegionAdmissionController(10, 1), + controlQueue: newControlQueue(), + tracker: newRegionTracker(), + storeAddr: "127.0.0.1:1", + } + + firstRegion := createFailureRecoveryTestRegion(t, 1, 1) + pendingRegion := createFailureRecoveryTestRegion(t, 2, 2) + submitRegionForAdmission(t, worker.admission, firstRegion, 100) + submitRegionForAdmission(t, worker.admission, pendingRegion, 100) + + ctx, cancel := context.WithCancel(context.Background()) + runErrCh := make(chan error, 1) + go func() { + runErrCh <- worker.Run(ctx) + }() + + require.Eventually(t, func() bool { + return errCacheLen(handler) == 2 + }, 5*time.Second, 10*time.Millisecond) + require.ElementsMatch(t, []uint64{1, 2}, snapshotErrCacheRegionIDs(handler)) + require.Equal(t, 0, worker.admission.stats().pending) + require.Equal(t, 0, worker.admission.stats().inflight) + + cancel() + require.ErrorIs(t, <-runErrCh, context.Canceled) } func TestProcessRegionSendTaskSendFailureCleansSentRequest(t *testing.T) { + admission := newTestRegionAdmissionController(10, 1) worker := ®ionRequestWorker{ - requestCache: newRequestCache(10), - store: &requestedStore{storeAddr: "store-1"}, - client: &subscriptionClient{}, + admission: admission, + controlQueue: newControlQueue(), + storeAddr: "store-1", + upstream: &upstreamHandle{}, + tracker: newRegionTracker(), } - worker.requestedRegions.subscriptions = make(map[SubscriptionID]regionFeedStates) - ctx := context.Background() region := prepareRegionForSendTest(createTestRegionInfo(1, 1)) - ok, err := worker.requestCache.add(ctx, region, false) - require.NoError(t, err) - require.True(t, ok) - require.Equal(t, 1, worker.requestCache.getPendingCount()) - - req, err := worker.requestCache.pop(ctx) - require.NoError(t, err) - worker.preFetchForConnecting = new(regionInfo) - *worker.preFetchForConnecting = req.regionInfo + req := admitRegionRequest(t, admission, region) + require.Equal(t, 1, admission.stats().inflight) sendErr := errors.New("send failed") conn := &ConnAndClient{ @@ -365,12 +571,46 @@ func TestProcessRegionSendTaskSendFailureCleansSentRequest(t *testing.T) { Conn: &grpc.ClientConn{}, } - err = worker.processRegionSendTask(ctx, conn) + err := worker.processRegionSendTask(t.Context(), conn, req) require.ErrorIs(t, err, sendErr) - require.Equal(t, 0, worker.requestCache.getPendingCount()) - require.Empty(t, worker.requestCache.sentRequests.regionReqs) - state := worker.getRegionState(req.regionInfo.subscribedSpan.subID, req.regionInfo.verID.GetID()) - require.True(t, state == nil || state.isStale(), "region state should be removed or marked stale after send failure") + require.Equal(t, 0, admission.stats().inflight) + state := worker.tracker.Get(req.regionInfo.subscribedSpan.subID, req.regionInfo.verID.GetID()) + require.NotNil(t, state) + require.True(t, state.isStale()) + var streamErr *storeStreamErr + require.ErrorAs(t, state.takeError(), &streamErr) +} + +func TestProcessRegionSendTaskDoesNotSendRemovedRequest(t *testing.T) { + admission := newTestRegionAdmissionController(1, 1) + worker := ®ionRequestWorker{ + admission: admission, + controlQueue: newControlQueue(), + storeAddr: "store-1", + upstream: &upstreamHandle{}, + tracker: newRegionTracker(), + } + region := prepareRegionForSendTest(createTestRegionInfo(1, 1)) + req := admitRegionRequest(t, admission, region) + require.True(t, req.abort()) + + ctx, cancel := context.WithCancel(context.Background()) + sendCh := make(chan *cdcpb.ChangeDataRequest, 1) + done := make(chan error, 1) + go func() { + done <- worker.processRegionSendTask(ctx, &ConnAndClient{ + Client: &mockEventFeedV2Client{sendCh: sendCh}, + Conn: &grpc.ClientConn{}, + }, req) + }() + + select { + case sentReq := <-sendCh: + t.Fatalf("removed request was sent: %+v", sentReq) + case <-time.After(50 * time.Millisecond): + } + cancel() + require.ErrorIs(t, <-done, context.Canceled) } func TestProcessRegionSendTaskSendEOFIsRetriable(t *testing.T) { @@ -390,37 +630,29 @@ func TestProcessRegionSendTaskSendEOFIsRetriable(t *testing.T) { for _, tc := range testCases { t.Run(tc.name, func(t *testing.T) { + admission := newTestRegionAdmissionController(10, 1) worker := ®ionRequestWorker{ - requestCache: newRequestCache(10), - store: &requestedStore{storeAddr: "store-1"}, - client: &subscriptionClient{}, + admission: admission, + controlQueue: newControlQueue(), + storeAddr: "store-1", + upstream: &upstreamHandle{}, + tracker: newRegionTracker(), } - worker.requestedRegions.subscriptions = make(map[SubscriptionID]regionFeedStates) - - ctx := context.Background() region := prepareRegionForSendTest(createTestRegionInfo(1, 1)) - ok, err := worker.requestCache.add(ctx, region, false) - require.NoError(t, err) - require.True(t, ok) - - req, err := worker.requestCache.pop(ctx) - require.NoError(t, err) - worker.preFetchForConnecting = new(regionInfo) - *worker.preFetchForConnecting = req.regionInfo + req := admitRegionRequest(t, admission, region) conn := &ConnAndClient{ Client: &mockEventFeedV2Client{sendErr: tc.sendErr}, Conn: &grpc.ClientConn{}, } - err = worker.processRegionSendTask(ctx, conn) + err := worker.processRegionSendTask(t.Context(), conn, req) var streamErr *storeStreamErr require.ErrorAs(t, err, &streamErr) - require.Equal(t, 0, worker.requestCache.getPendingCount()) - require.Empty(t, worker.requestCache.sentRequests.regionReqs) + require.Equal(t, 0, admission.stats().inflight) - state := worker.getRegionState(req.regionInfo.subscribedSpan.subID, req.regionInfo.verID.GetID()) + state := worker.tracker.Get(req.regionInfo.subscribedSpan.subID, req.regionInfo.verID.GetID()) require.NotNil(t, state) require.True(t, state.isStale()) @@ -430,6 +662,42 @@ func TestProcessRegionSendTaskSendEOFIsRetriable(t *testing.T) { } } +func TestProcessRegionSendTaskHandlesDeregisterFromControlQueue(t *testing.T) { + ds := &mockRegionEventDynamicStream{} + worker := ®ionRequestWorker{ + admission: newTestRegionAdmissionController(1, 1), + controlQueue: newControlQueue(), + storeAddr: "store-1", + upstream: &upstreamHandle{clusterID: 42}, + eventSink: ®ionEventSink{ds: ds}, + tracker: newRegionTracker(), + } + state := ®ionFeedState{worker: worker} + require.True(t, worker.tracker.Add(1, 1, state)) + worker.controlQueue.push(deregisterRequest{subID: 1, filterLoop: true}) + + ctx, cancel := context.WithCancel(context.Background()) + sendCh := make(chan *cdcpb.ChangeDataRequest, 1) + done := make(chan error, 1) + go func() { + done <- worker.processRegionSendTask(ctx, &ConnAndClient{ + Client: &mockEventFeedV2Client{sendCh: sendCh}, + Conn: &grpc.ClientConn{}, + }, nil) + }() + + req := <-sendCh + require.Equal(t, uint64(42), req.Header.ClusterId) + require.Equal(t, uint64(1), req.RequestId) + require.True(t, req.FilterLoop) + require.NotNil(t, req.GetDeregister()) + require.Eventually(t, func() bool { + return worker.tracker.Get(1, 1) == nil + }, time.Second, 10*time.Millisecond) + cancel() + require.ErrorIs(t, <-done, context.Canceled) +} + func TestReceiveAndDispatchChangeEventsEOFIsRetriable(t *testing.T) { testCases := []struct { name string @@ -447,9 +715,7 @@ func TestReceiveAndDispatchChangeEventsEOFIsRetriable(t *testing.T) { for _, tc := range testCases { t.Run(tc.name, func(t *testing.T) { - worker := ®ionRequestWorker{ - store: &requestedStore{storeAddr: "store-1"}, - } + worker := ®ionRequestWorker{storeAddr: "store-1"} conn := &ConnAndClient{ Client: &mockEventFeedV2Client{recvErr: tc.recvErr}, Conn: &grpc.ClientConn{}, diff --git a/logservice/logpuller/region_state.go b/logservice/logpuller/region_state.go index 40e98bdb26..422c3bb39c 100644 --- a/logservice/logpuller/region_state.go +++ b/logservice/logpuller/region_state.go @@ -15,6 +15,7 @@ package logpuller import ( "sync" + "sync/atomic" "github.com/pingcap/kvproto/pkg/cdcpb" "github.com/pingcap/ticdc/heartbeatpb" @@ -52,11 +53,6 @@ type regionInfo struct { scanPriority cdcpb.ScanPriority } -func (s *regionInfo) isStopped() bool { - // lockedRange only nil when the region's subscribedTable is stopped. - return s.lockedRangeState == nil -} - func newRegionInfo( verID tikv.RegionVerID, span heartbeatpb.TableSpan, @@ -70,7 +66,7 @@ func newRegionInfo( rpcCtx: rpcCtx, subscribedSpan: subscribedSpan, filterLoop: filterLoop, - scanPriority: TaskLowPrior.scanPriority(), + scanPriority: cdcpb.ScanPriority_SCAN_PRIORITY_LOW, } } @@ -94,6 +90,9 @@ type regionFeedState struct { region regionInfo requestID uint64 // It is also the subscription ID matcher *matcher + // onInitialized runs once when the region finishes its first successful + // initialization for this request lifecycle. + onInitialized func(*regionFeedState) // Transform: normal -> stopped -> removed. // normal: the region is in replicating. @@ -107,20 +106,27 @@ type regionFeedState struct { // `err` is used to retrieve errors generated outside. err error } + regionReq atomic.Pointer[regionReq] worker *regionRequestWorker } -func newRegionFeedState(region regionInfo, requestID uint64, worker *regionRequestWorker) *regionFeedState { - return ®ionFeedState{ - region: region, - requestID: requestID, - worker: worker, +func newRegionFeedState( + region regionInfo, + requestID uint64, + worker *regionRequestWorker, + request *regionReq, + onInitialized func(*regionFeedState), +) *regionFeedState { + state := ®ionFeedState{ + region: region, + requestID: requestID, + matcher: newMatcher(), + onInitialized: onInitialized, + worker: worker, } -} - -func (s *regionFeedState) start() { - s.matcher = newMatcher() + state.regionReq.Store(request) + return state } // mark regionFeedState as stopped with the given error if possible. @@ -131,7 +137,7 @@ func (s *regionFeedState) markStopped(err error) { s.state.v = stateStopped s.state.err = err } - s.worker.requestCache.markStopped(s.region.subscribedSpan.subID, s.region.verID.GetID()) + s.abortScanIfNeeded() } // mark regionFeedState as removed if possible. @@ -143,7 +149,7 @@ func (s *regionFeedState) markRemoved() (changed bool) { changed = true s.matcher.clear() } - s.worker.requestCache.markStopped(s.region.subscribedSpan.subID, s.region.verID.GetID()) + s.abortScanIfNeeded() return } @@ -166,8 +172,27 @@ func (s *regionFeedState) isInitialized() bool { } func (s *regionFeedState) setInitialized() { - s.region.lockedRangeState.Initialized.Store(true) - s.worker.requestCache.resolve(s.region.subscribedSpan.subID, s.region.verID.GetID()) + if !s.region.lockedRangeState.Initialized.CompareAndSwap(false, true) { + return + } + s.finishScan() + if s.onInitialized != nil { + s.onInitialized(s) + } +} + +func (s *regionFeedState) finishScan() { + request := s.regionReq.Swap(nil) + if request != nil { + request.finish() + } +} + +func (s *regionFeedState) abortScanIfNeeded() { + request := s.regionReq.Swap(nil) + if request != nil { + request.abort() + } } func (s *regionFeedState) getRegionID() uint64 { diff --git a/logservice/logpuller/region_tracker.go b/logservice/logpuller/region_tracker.go new file mode 100644 index 0000000000..81a27dc86c --- /dev/null +++ b/logservice/logpuller/region_tracker.go @@ -0,0 +1,124 @@ +// Copyright 2026 PingCAP, Inc. +// +// 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, +// See the License for the specific language governing permissions and +// limitations under the License. + +package logpuller + +import ( + "maps" + "slices" + "sync" +) + +type regionStatesByID map[uint64]*regionFeedState + +// regionTracker owns the region states tracked by one region request worker. +type regionTracker struct { + mu sync.RWMutex + + statesBySubscription map[SubscriptionID]regionStatesByID +} + +func newRegionTracker() *regionTracker { + return ®ionTracker{ + statesBySubscription: make(map[SubscriptionID]regionStatesByID), + } +} + +// Get returns the state tracked by a subscription and region. +func (t *regionTracker) Get(subscriptionID SubscriptionID, regionID uint64) *regionFeedState { + t.mu.RLock() + defer t.mu.RUnlock() + + if states, ok := t.statesBySubscription[subscriptionID]; ok { + return states[regionID] + } + return nil +} + +// Add records a region state unless the same subscription and region is +// already tracked. +func (t *regionTracker) Add( + subscriptionID SubscriptionID, + regionID uint64, + state *regionFeedState, +) bool { + t.mu.Lock() + defer t.mu.Unlock() + + states := t.statesBySubscription[subscriptionID] + if states == nil { + states = make(regionStatesByID) + t.statesBySubscription[subscriptionID] = states + } + if _, ok := states[regionID]; ok { + return false + } + states[regionID] = state + return true +} + +// RemoveIf removes a region only when it is still tracked by the expected +// state. It prevents delayed events for an old state from removing a newer +// owner. +func (t *regionTracker) RemoveIf( + subscriptionID SubscriptionID, + regionID uint64, + expected *regionFeedState, +) bool { + if expected == nil { + return false + } + + t.mu.Lock() + defer t.mu.Unlock() + + if states, ok := t.statesBySubscription[subscriptionID]; ok { + if states[regionID] != expected { + return false + } + delete(states, regionID) + if len(states) == 0 { + delete(t.statesBySubscription, subscriptionID) + } + return true + } + return false +} + +// TakeSubscription removes and returns all states tracked by a subscription. +func (t *regionTracker) TakeSubscription(subscriptionID SubscriptionID) []*regionFeedState { + t.mu.Lock() + states := t.statesBySubscription[subscriptionID] + delete(t.statesBySubscription, subscriptionID) + t.mu.Unlock() + + return slices.Collect(maps.Values(states)) +} + +// Drain removes and returns all tracked states. +func (t *regionTracker) Drain() []*regionFeedState { + t.mu.Lock() + statesBySubscription := t.statesBySubscription + t.statesBySubscription = make(map[SubscriptionID]regionStatesByID) + t.mu.Unlock() + + totalStates := 0 + for _, states := range statesBySubscription { + totalStates += len(states) + } + drainedStates := make([]*regionFeedState, 0, totalStates) + for _, states := range statesBySubscription { + drainedStates = append(drainedStates, slices.Collect(maps.Values(states))...) + } + return drainedStates +} diff --git a/logservice/logpuller/region_tracker_test.go b/logservice/logpuller/region_tracker_test.go new file mode 100644 index 0000000000..8a08ce0d94 --- /dev/null +++ b/logservice/logpuller/region_tracker_test.go @@ -0,0 +1,58 @@ +// Copyright 2026 PingCAP, Inc. +// +// 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, +// See the License for the specific language governing permissions and +// limitations under the License. + +package logpuller + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestRegionTrackerOperations(t *testing.T) { + tracker := newRegionTracker() + state1 := ®ionFeedState{} + state2 := ®ionFeedState{} + state3 := ®ionFeedState{} + + require.Nil(t, tracker.Get(1, 1)) + require.True(t, tracker.Add(1, 1, state1)) + require.Same(t, state1, tracker.Get(1, 1)) + require.True(t, tracker.Add(1, 2, state2)) + require.True(t, tracker.Add(2, 3, state3)) + + require.ElementsMatch(t, []*regionFeedState{state1, state2}, tracker.TakeSubscription(1)) + require.Nil(t, tracker.Get(1, 1)) + require.Nil(t, tracker.Get(1, 2)) + require.Empty(t, tracker.TakeSubscription(1)) + + drained := tracker.Drain() + require.ElementsMatch(t, []*regionFeedState{state3}, drained) + require.Nil(t, tracker.Get(2, 3)) + require.Empty(t, tracker.Drain()) +} + +func TestRegionTrackerAddRejectsDuplicate(t *testing.T) { + tracker := newRegionTracker() + oldState := ®ionFeedState{} + newState := ®ionFeedState{} + + require.True(t, tracker.Add(1, 1, oldState)) + require.False(t, tracker.Add(1, 1, newState)) + + require.False(t, tracker.RemoveIf(1, 1, newState)) + require.Same(t, oldState, tracker.Get(1, 1)) + require.True(t, tracker.RemoveIf(1, 1, oldState)) + require.Nil(t, tracker.Get(1, 1)) + require.False(t, tracker.RemoveIf(1, 1, oldState)) +} diff --git a/logservice/logpuller/scan_priority.go b/logservice/logpuller/scan_priority.go index 7f8480265a..176b9122d5 100644 --- a/logservice/logpuller/scan_priority.go +++ b/logservice/logpuller/scan_priority.go @@ -18,6 +18,7 @@ import ( "sync/atomic" "time" + "github.com/pingcap/kvproto/pkg/cdcpb" "github.com/pingcap/ticdc/pkg/pdutil" "github.com/tikv/client-go/v2/oracle" ) @@ -48,14 +49,14 @@ func (p *scanPriorityPolicy) observeSpanResolved(resolvedTs uint64) bool { // resolve returns the effective priority after combining inherited, span, and // region progress. A high priority decision is never downgraded. func (p *scanPriorityPolicy) resolve( - inherited TaskType, + inherited cdcpb.ScanPriority, regionResolvedTs uint64, currentTime time.Time, -) TaskType { - if inherited == TaskHighPrior || p.everCaughtUp.Load() || p.isTsClose(regionResolvedTs, currentTime) { - return TaskHighPrior +) cdcpb.ScanPriority { + if isHighScanPriority(inherited) || p.everCaughtUp.Load() || p.isTsClose(regionResolvedTs, currentTime) { + return cdcpb.ScanPriority_SCAN_PRIORITY_HIGH } - return TaskLowPrior + return cdcpb.ScanPriority_SCAN_PRIORITY_LOW } func (p *scanPriorityPolicy) isTsClose(ts uint64, currentTime time.Time) bool { diff --git a/logservice/logpuller/scan_priority_test.go b/logservice/logpuller/scan_priority_test.go index d5501efcd2..629069cbf0 100644 --- a/logservice/logpuller/scan_priority_test.go +++ b/logservice/logpuller/scan_priority_test.go @@ -37,45 +37,45 @@ func TestScanPriorityPolicyResolve(t *testing.T) { for _, tc := range []struct { name string - inherited TaskType + inherited cdcpb.ScanPriority regionResolvedTs uint64 - expected TaskType + expected cdcpb.ScanPriority }{ { name: "zero resolved ts", - inherited: TaskLowPrior, + inherited: cdcpb.ScanPriority_SCAN_PRIORITY_LOW, regionResolvedTs: 0, - expected: TaskLowPrior, + expected: cdcpb.ScanPriority_SCAN_PRIORITY_LOW, }, { name: "recent region", - inherited: TaskLowPrior, + inherited: cdcpb.ScanPriority_SCAN_PRIORITY_LOW, regionResolvedTs: oracle.GoTimeToTS(currentTime.Add(-29 * time.Minute)), - expected: TaskHighPrior, + expected: cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, }, { name: "threshold boundary", - inherited: TaskLowPrior, + inherited: cdcpb.ScanPriority_SCAN_PRIORITY_LOW, regionResolvedTs: oracle.GoTimeToTS(currentTime.Add(-30 * time.Minute)), - expected: TaskHighPrior, + expected: cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, }, { name: "old region", - inherited: TaskLowPrior, + inherited: cdcpb.ScanPriority_SCAN_PRIORITY_LOW, regionResolvedTs: oracle.GoTimeToTS(currentTime.Add(-31 * time.Minute)), - expected: TaskLowPrior, + expected: cdcpb.ScanPriority_SCAN_PRIORITY_LOW, }, { name: "future region", - inherited: TaskLowPrior, + inherited: cdcpb.ScanPriority_SCAN_PRIORITY_LOW, regionResolvedTs: oracle.GoTimeToTS(currentTime.Add(time.Minute)), - expected: TaskHighPrior, + expected: cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, }, { name: "inherited high", - inherited: TaskHighPrior, + inherited: cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, regionResolvedTs: oracle.GoTimeToTS(currentTime.Add(-time.Hour)), - expected: TaskHighPrior, + expected: cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, }, } { t.Run(tc.name, func(t *testing.T) { @@ -93,8 +93,8 @@ func TestScanPriorityPolicyRemainsHighAfterCatchUp(t *testing.T) { require.False(t, policy.observeSpanResolved(oracle.GoTimeToTS(currentTime.Add(-31*time.Minute)))) require.True(t, policy.observeSpanResolved(oracle.GoTimeToTS(currentTime.Add(-time.Minute)))) require.False(t, policy.observeSpanResolved(oracle.GoTimeToTS(currentTime))) - require.Equal(t, TaskHighPrior, policy.resolve( - TaskLowPrior, + require.Equal(t, cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, policy.resolve( + cdcpb.ScanPriority_SCAN_PRIORITY_LOW, oracle.GoTimeToTS(currentTime.Add(-time.Hour)), currentTime, )) @@ -106,8 +106,10 @@ func TestScanPriorityUsesRestoredRegionProgress(t *testing.T) { pdClock := pdutil.NewClock4Test() pdClock.(*pdutil.Clock4Test).SetTS(currentTs) client := &subscriptionClient{ - pdClock: pdClock, - regionTaskQueue: priorityqueue.New[PriorityTask](), + upstream: &upstreamHandle{pdClock: pdClock}, + regionScheduler: ®ionRequestScheduler{ + taskQueue: priorityqueue.New[*regionPriorityTask](), + }, } startTs := oracle.GoTimeToTS(currentTime.Add(-time.Hour)) @@ -125,11 +127,11 @@ func TestScanPriorityUsesRestoredRegionProgress(t *testing.T) { } region := newRegionInfo(tikv.NewRegionVerID(1, 1, 1), rawSpan, nil, span, false) - client.scheduleRegionRequest(context.Background(), region, TaskLowPrior) - firstTask := popRegionPriorityTask(t, client.regionTaskQueue) - require.Equal(t, TaskLowPrior, firstTask.taskType) + client.scheduleRegionRequest(context.Background(), region) + firstTask := popRegionPriorityTask(t, client.regionScheduler.taskQueue) + require.Equal(t, cdcpb.ScanPriority_SCAN_PRIORITY_LOW, firstTask.priority()) - firstRegion := firstTask.GetRegionInfo() + firstRegion := firstTask.regionInfo firstRegion.lockedRangeState.ResolvedTs.Store(oracle.GoTimeToTS(currentTime.Add(-time.Minute))) span.rangeLock.UnlockRange( firstRegion.span.StartKey, @@ -139,21 +141,21 @@ func TestScanPriorityUsesRestoredRegionProgress(t *testing.T) { ) retryRegion := newRegionInfo(tikv.NewRegionVerID(1, 1, 2), rawSpan, nil, span, false) - client.scheduleRegionRequest(context.Background(), retryRegion, TaskLowPrior) - retryTask := popRegionPriorityTask(t, client.regionTaskQueue) - require.Equal(t, TaskHighPrior, retryTask.taskType) - require.Equal(t, cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, retryTask.GetRegionInfo().scanPriority) + client.scheduleRegionRequest(context.Background(), retryRegion) + retryTask := popRegionPriorityTask(t, client.regionScheduler.taskQueue) + require.Equal(t, cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, retryTask.priority()) + require.Equal(t, cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, retryTask.regionInfo.scanPriority) require.False(t, span.priorityPolicy.everCaughtUp.Load()) } func popRegionPriorityTask( t *testing.T, - queue *priorityqueue.PriorityQueue[PriorityTask], + queue *priorityqueue.PriorityQueue[*regionPriorityTask], ) *regionPriorityTask { t.Helper() ctx, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() task, err := queue.Pop(ctx) require.NoError(t, err) - return task.(*regionPriorityTask) + return task } diff --git a/logservice/logpuller/subscription_client.go b/logservice/logpuller/subscription_client.go index fc3a0725c3..af039a9ebd 100644 --- a/logservice/logpuller/subscription_client.go +++ b/logservice/logpuller/subscription_client.go @@ -15,10 +15,10 @@ package logpuller import ( "context" - "sync" "sync/atomic" "time" + "github.com/pingcap/kvproto/pkg/cdcpb" "github.com/pingcap/kvproto/pkg/metapb" "github.com/pingcap/log" "github.com/pingcap/ticdc/heartbeatpb" @@ -27,20 +27,14 @@ import ( "github.com/pingcap/ticdc/pkg/common" appcontext "github.com/pingcap/ticdc/pkg/common/context" "github.com/pingcap/ticdc/pkg/config" - "github.com/pingcap/ticdc/pkg/errors" "github.com/pingcap/ticdc/pkg/metrics" "github.com/pingcap/ticdc/pkg/pdutil" "github.com/pingcap/ticdc/pkg/security" "github.com/pingcap/ticdc/pkg/spanz" "github.com/pingcap/ticdc/pkg/util" - "github.com/pingcap/ticdc/utils/priorityqueue" - "github.com/prometheus/client_golang/prometheus" - kvclientv2 "github.com/tikv/client-go/v2/kv" - "github.com/tikv/client-go/v2/oracle" "github.com/tikv/client-go/v2/tikv" pd "github.com/tikv/pd/client" "go.uber.org/zap" - "go.uber.org/zap/zapcore" "golang.org/x/sync/errgroup" ) @@ -89,16 +83,23 @@ type rangeTask struct { span heartbeatpb.TableSpan subscribedSpan *subscribedSpan filterLoop bool - priority TaskType + priority cdcpb.ScanPriority } -type SubscriptionClientConfig struct { - // The number of region request workers to send region task for every tikv store - RegionRequestWorkerPerStore uint +// upstreamHandle contains the stable TiKV and PD dependencies shared by the +// region request pipeline. +type upstreamHandle struct { + pd pd.Client + regionCache *tikv.RegionCache + pdClock pdutil.Clock + credential *security.Credential + clusterID uint64 } -type sharedClientMetrics struct { - batchResolvedSize prometheus.Observer +// initialize loads the cluster metadata needed by region request workers. It +// must run before the scheduler starts any workers. +func (u *upstreamHandle) initialize(ctx context.Context) { + u.clusterID = u.pd.GetClusterID(ctx) } // subscriptionClient is used to subscribe events of table ranges from TiKV. @@ -122,21 +123,11 @@ type SubscriptionClient interface { } type subscriptionClient struct { - ctx context.Context - cancel context.CancelFunc - config *SubscriptionClientConfig - metrics sharedClientMetrics - clusterID uint64 - - pd pd.Client - regionCache *tikv.RegionCache - pdClock pdutil.Clock - lockResolver txnutil.LockResolver - - stores sync.Map + ctx context.Context + cancel context.CancelFunc + upstream *upstreamHandle - // the credential to connect tikv - credential *security.Credential + lockResolver txnutil.LockResolver // failureHandler handles failed regions and owns reschedule/retry decisions. failureHandler *regionFailureHandler @@ -144,13 +135,14 @@ type subscriptionClient struct { eventSink *regionEventSink // spanRegistry tracks subscribed spans and owns span-level background tasks. spanRegistry *spanRegistry + // regionScheduler assigns locked region requests to per-store workers. + regionScheduler *regionRequestScheduler + // memoryQuota owns event-memory accounting and initial-scan admission. + memoryQuota *memoryQuotaController // rangeTaskCh is used to receive range tasks. // The tasks will be handled in `handleRangeTask` goroutine. rangeTaskCh chan rangeTask - // regionTaskQueue is used to receive region tasks with priority. - // The region will be handled in `handleRegions` goroutine. - regionTaskQueue *priorityqueue.PriorityQueue[PriorityTask] // resolveLockTaskCh is used to receive resolve lock tasks. // The tasks will be handled in `handleResolveLockTasks` goroutine. resolveLockTaskCh chan resolveLockTask @@ -159,33 +151,45 @@ type subscriptionClient struct { // NewSubscriptionClient creates a client. func NewSubscriptionClient( - config *SubscriptionClientConfig, pd pd.Client, lockResolver txnutil.LockResolver, credential *security.Credential, ) SubscriptionClient { subClient := &subscriptionClient{ - config: config, - - stores: sync.Map{}, - pd: pd, - regionCache: appcontext.GetService[*tikv.RegionCache](appcontext.RegionCache), - pdClock: appcontext.GetService[pdutil.Clock](appcontext.DefaultPDClock), + upstream: &upstreamHandle{ + pd: pd, + regionCache: appcontext.GetService[*tikv.RegionCache](appcontext.RegionCache), + pdClock: appcontext.GetService[pdutil.Clock](appcontext.DefaultPDClock), + credential: credential, + }, lockResolver: lockResolver, - credential: credential, - rangeTaskCh: make(chan rangeTask, 1024), - regionTaskQueue: priorityqueue.New[PriorityTask](), resolveLockTaskCh: make(chan resolveLockTask, 1024), resolveLockRateLimiter: newResolveLockRateLimiter(), } subClient.ctx, subClient.cancel = context.WithCancel(context.Background()) - subClient.failureHandler = newRegionFailureHandler(subClient) - subClient.eventSink = newRegionEventSink(subClient.failureHandler) - subClient.spanRegistry = newSpanRegistry(subClient.pd, subClient.pdClock) - - subClient.initMetrics() + pullerConfig := config.GetGlobalServerConfig().Debug.Puller + subClient.memoryQuota = newMemoryQuotaController( + pullerConfig.MemoryQuota, pullerConfig.ScanBaseSize) + subClient.failureHandler = newRegionFailureHandler( + subClient.upstream.regionCache, + subClient.onTableDrained, + subClient.scheduleRegionRequest, + subClient.scheduleRangeRequest, + ) + subClient.eventSink = newRegionEventSink( + subClient.ctx, + subClient.failureHandler, + subClient.memoryQuota, + ) + subClient.spanRegistry = newSpanRegistry(subClient.upstream.pd, subClient.upstream.pdClock) + subClient.regionScheduler = newRegionRequestScheduler( + subClient.upstream, + subClient.eventSink, + subClient.failureHandler, + subClient.memoryQuota, + ) return subClient } @@ -198,11 +202,6 @@ func (s *subscriptionClient) AllocSubscriptionID() SubscriptionID { return SubscriptionID(subscriptionIDGen.Add(1)) } -func (s *subscriptionClient) initMetrics() { - // TODO: fix metrics - s.metrics.batchResolvedSize = metrics.BatchResolvedEventSize.WithLabelValues("event-store") -} - func (s *subscriptionClient) updateMetrics(ctx context.Context) error { ticker := time.NewTicker(10 * time.Second) defer ticker.Stop() @@ -211,20 +210,9 @@ func (s *subscriptionClient) updateMetrics(ctx context.Context) error { case <-ctx.Done(): return ctx.Err() case <-ticker.C: - pendingRegionReqCount := 0 - s.stores.Range(func(_, value any) bool { - store := value.(*requestedStore) - store.requestWorkers.RLock() - for _, worker := range store.requestWorkers.s { - worker.requestCache.clearStaleRequest() - pendingRegionReqCount += worker.requestCache.getPendingCount() - } - store.requestWorkers.RUnlock() - return true - }) - - metrics.SubscriptionClientRequestedRegionCount.WithLabelValues("pending").Set(float64(pendingRegionReqCount)) + s.regionScheduler.UpdateMetrics() s.eventSink.UpdateMetrics() + s.memoryQuota.UpdateMetrics() s.spanRegistry.UpdateMetrics() } } @@ -260,7 +248,7 @@ func (s *subscriptionClient) Subscribe( advanceResolvedTs, advanceInterval, bdrMode, - s.pdClock, + s.upstream.pdClock, time.Duration(config.GetGlobalServerConfig().Debug.Puller.OldStartTsScanLowPriorityThreshold), ) s.spanRegistry.Add(rt) @@ -269,9 +257,13 @@ func (s *subscriptionClient) Subscribe( select { case <-s.ctx.Done(): log.Warn("subscribes span failed, the subscription client has closed") - case s.rangeTaskCh <- rangeTask{span: span, subscribedSpan: rt, filterLoop: rt.filterLoop, priority: TaskLowPrior}: - log.Info("subscribes span done", - zap.Uint64("subscriptionID", uint64(subID)), + case s.rangeTaskCh <- rangeTask{ + span: span, + subscribedSpan: rt, + filterLoop: rt.filterLoop, + priority: cdcpb.ScanPriority_SCAN_PRIORITY_LOW, + }: + log.Info("subscribes span done", zap.Uint64("subscriptionID", uint64(subID)), zap.Int64("tableID", span.TableID), zap.Uint64("startTs", startTs), zap.String("startKey", spanz.HexKey(span.StartKey)), zap.String("endKey", spanz.HexKey(span.EndKey))) } @@ -293,27 +285,19 @@ func (s *subscriptionClient) Unsubscribe(subID SubscriptionID) { zap.Bool("exists", rt != nil)) } -func (s *subscriptionClient) pushRegionEventToDS(subID SubscriptionID, event regionEvent) { - s.eventSink.Push(subID, event) -} - func (s *subscriptionClient) Run(ctx context.Context) error { - // s.consume = consume - if s.pd == nil { - log.Warn("subscription client should be in test mode, skip run") - return nil - } - s.clusterID = s.pd.GetClusterID(ctx) + s.upstream.initialize(ctx) g, ctx := errgroup.WithContext(ctx) - g.Go(func() error { return s.updateMetrics(ctx) }) - g.Go(func() error { return s.eventSink.Run(ctx) }) + // The goroutines are listed by data flow; errgroup does not guarantee their + // actual startup order. g.Go(func() error { return s.handleRangeTasks(ctx) }) - g.Go(func() error { return s.handleRegions(ctx, g) }) + g.Go(func() error { return s.regionScheduler.Run(ctx, g) }) g.Go(func() error { return s.failureHandler.Run(ctx) }) - g.Go(func() error { return s.handleResolveLockTasks(ctx) }) g.Go(func() error { return s.spanRegistry.Run(ctx) }) + g.Go(func() error { return s.handleResolveLockTasks(ctx) }) + g.Go(func() error { return s.updateMetrics(ctx) }) log.Info("subscription client starts") defer log.Info("subscription client exits") @@ -324,7 +308,7 @@ func (s *subscriptionClient) Run(ctx context.Context) error { func (s *subscriptionClient) Close(ctx context.Context) error { s.cancel() s.eventSink.Close() - s.regionTaskQueue.Close() + s.regionScheduler.Close() return nil } @@ -332,11 +316,12 @@ func (s *subscriptionClient) setTableStopped(rt *subscribedSpan) { log.Info("subscription client starts to stop table", zap.Uint64("subscriptionID", uint64(rt.subID))) - // Set stopped to true so we can stop handling region events from the table. - // Then send a special singleRegionInfo to regionRouter to deregister the table - // from all TiKV instances. + // Set stopped to true so we can stop handling region events from the table, + // then notify every existing worker to deregister the subscription. if rt.stopped.CompareAndSwap(false, true) { - s.regionTaskQueue.Push(NewRegionPriorityTask(TaskHighPrior, regionInfo{subscribedSpan: rt, filterLoop: rt.filterLoop}, s.pdClock.CurrentTS())) + // Wake event receivers and scan admission so they can observe stopped. + s.memoryQuota.WakeAll() + s.regionScheduler.BroadcastDeregister(rt.subID, rt.filterLoop) if rt.rangeLock.Stop() { s.onTableDrained(rt) } @@ -356,182 +341,6 @@ func (s *subscriptionClient) onTableDrained(rt *subscribedSpan) { s.spanRegistry.Remove(rt.subID) } -// Note: don't block the caller, otherwise there may be deadlock -func (s *subscriptionClient) onRegionFail(errInfo regionErrorInfo) { - s.failureHandler.Report(errInfo) -} - -// requestedStore represents a store that has been connected. -type requestedStore struct { - storeAddr string - // Use to select a worker to send request. - nextWorker atomic.Uint32 - - requestWorkers struct { - sync.RWMutex - s []*regionRequestWorker - } -} - -func (rs *requestedStore) getRequestWorker() *regionRequestWorker { - rs.requestWorkers.RLock() - defer rs.requestWorkers.RUnlock() - - index := rs.nextWorker.Add(1) % uint32(len(rs.requestWorkers.s)) - return rs.requestWorkers.s[index] -} - -// handleRegions receives regionInfo from regionTaskQueue and attach rpcCtx to them, -// then send them to corresponding requestedStore. -func (s *subscriptionClient) handleRegions(ctx context.Context, eg *errgroup.Group) error { - cfg := config.GetGlobalServerConfig() - pendingRegionRequestQueueSize := cfg.Debug.Puller.PendingRegionRequestQueueSize - getStore := func(storeAddr string) *requestedStore { - var rs *requestedStore - if v, ok := s.stores.Load(storeAddr); ok { - rs = v.(*requestedStore) - return rs - } - - rs = &requestedStore{storeAddr: storeAddr} - rs.requestWorkers.s = make([]*regionRequestWorker, 0, s.config.RegionRequestWorkerPerStore) - s.stores.Store(storeAddr, rs) - - perWorkerQueueSize := pendingRegionRequestQueueSize / int(s.config.RegionRequestWorkerPerStore) - if perWorkerQueueSize <= 0 { - log.Warn("pending region request queue size is smaller than the number of workers, adjust per worker queue size to 1", - zap.Int("pendingRegionRequestQueueSize", pendingRegionRequestQueueSize), - zap.Uint("regionRequestWorkerPerStore", s.config.RegionRequestWorkerPerStore)) - perWorkerQueueSize = 1 - } - - rs.requestWorkers.Lock() - for i := uint(0); i < s.config.RegionRequestWorkerPerStore; i++ { - requestWorker := newRegionRequestWorker(ctx, s, s.credential, eg, rs, perWorkerQueueSize) - rs.requestWorkers.s = append(rs.requestWorkers.s, requestWorker) - } - rs.requestWorkers.Unlock() - return rs - } - - defer func() { - s.stores.Range(func(_, value any) bool { - rs := value.(*requestedStore) - - rs.requestWorkers.RLock() - for _, w := range rs.requestWorkers.s { - w.requestCache.clear() - } - rs.requestWorkers.RUnlock() - - return true - }) - }() - - for { - select { - case <-ctx.Done(): - return ctx.Err() - default: - } - // Use blocking Pop to wait for tasks - regionTask, err := s.regionTaskQueue.Pop(ctx) - if err != nil { - if errors.Is(err, priorityqueue.ErrClosed) { - return nil - } - return err - } - - region := regionTask.GetRegionInfo() - if region.isStopped() { - enqueued, err := s.enqueueRegionToAllStores(ctx, region) - if err != nil { - return err - } - if !enqueued { - log.Debug("enqueue stop request failed, retry later", - zap.Uint64("subscriptionID", uint64(region.subscribedSpan.subID))) - s.regionTaskQueue.Push(regionTask) - } - continue - } - - region, ok := s.attachRPCContextForRegion(ctx, region) - // If attachRPCContextForRegion fails, the region will be re-scheduled. - if !ok { - continue - } - - store := getStore(region.rpcCtx.Addr) - worker := store.getRequestWorker() - force := regionTask.Priority() <= forcedPriorityBase - - ok, err = worker.add(ctx, region, force) - if err != nil { - log.Warn("subscription client add region request failed", - zap.Uint64("subscriptionID", uint64(region.subscribedSpan.subID)), - zap.Uint64("regionID", region.verID.GetID()), - zap.Error(err)) - return err - } - - if !ok { - s.regionTaskQueue.Push(regionTask) - continue - } - - log.Debug("subscription client will request a region", - zap.Uint64("workID", worker.workerID), - zap.Uint64("subscriptionID", uint64(region.subscribedSpan.subID)), - zap.Uint64("regionID", region.verID.GetID()), - zap.String("addr", store.storeAddr)) - } -} - -func (s *subscriptionClient) enqueueRegionToAllStores(ctx context.Context, region regionInfo) (bool, error) { - enqueued := true - var firstErr error - s.stores.Range(func(_ any, value any) bool { - rs := value.(*requestedStore) - rs.requestWorkers.RLock() - workers := rs.requestWorkers.s - rs.requestWorkers.RUnlock() - for _, worker := range workers { - ok, err := worker.add(ctx, region, true) - if err != nil { - firstErr = err - enqueued = false - return false - } - if !ok { - enqueued = false - // It is likely the store is busy, no need to try other workers in this store now. - break - } - } - return true - }) - return enqueued, firstErr -} - -func (s *subscriptionClient) attachRPCContextForRegion(ctx context.Context, region regionInfo) (regionInfo, bool) { - bo := tikv.NewBackoffer(ctx, tikvRequestMaxBackoff) - rpcCtx, err := s.regionCache.GetTiKVRPCContext(bo, region.verID, kvclientv2.ReplicaReadLeader, 0) - if rpcCtx != nil { - region.rpcCtx = rpcCtx - return region, true - } - if err != nil { - log.Debug("subscription client get rpc context fail", - zap.Uint64("subscriptionID", uint64(region.subscribedSpan.subID)), - zap.Uint64("regionID", region.verID.GetID()), - zap.Error(err)) - } - s.onRegionFail(newRegionErrorInfo(region, &rpcCtxUnavailableErr{verID: region.verID})) - return region, false -} - func (s *subscriptionClient) handleRangeTasks(ctx context.Context) error { g, ctx := errgroup.WithContext(ctx) // Limit the concurrent number of goroutines to convert range tasks to region tasks. @@ -542,7 +351,7 @@ func (s *subscriptionClient) handleRangeTasks(ctx context.Context) error { return ctx.Err() case task := <-s.rangeTaskCh: g.Go(func() error { - return s.divideSpanAndScheduleRegionRequests(ctx, task.span, task.subscribedSpan, task.filterLoop, task.priority) + return s.divideSpanAndScheduleRegionRequests(ctx, task) }) } } @@ -555,11 +364,11 @@ func (s *subscriptionClient) handleRangeTasks(ctx context.Context) error { // 3. Schedule a region request to subscribe the region. func (s *subscriptionClient) divideSpanAndScheduleRegionRequests( ctx context.Context, - span heartbeatpb.TableSpan, - subscribedSpan *subscribedSpan, - filterLoop bool, - taskType TaskType, + task rangeTask, ) error { + span := task.span + subscribedSpan := task.subscribedSpan + // Limit the number of regions loaded at a time to make the load more stable. limit := 1024 nextSpan := span @@ -576,7 +385,8 @@ func (s *subscriptionClient) divideSpanAndScheduleRegionRequests( zap.Any("span", common.FormatTableSpan(&nextSpan))) backoff := tikv.NewBackoffer(ctx, tikvRequestMaxBackoff) - regions, err := s.regionCache.BatchLoadRegionsWithKeyRange(backoff, nextSpan.StartKey, nextSpan.EndKey, limit) + regions, err := s.upstream.regionCache.BatchLoadRegionsWithKeyRange( + backoff, nextSpan.StartKey, nextSpan.EndKey, limit) if err != nil { log.Warn("subscription client load regions failed", zap.Uint64("subscriptionID", uint64(subscribedSpan.subID)), @@ -619,10 +429,11 @@ func (s *subscriptionClient) divideSpanAndScheduleRegionRequests( } verID := tikv.NewRegionVerID(regionMeta.Id, regionMeta.RegionEpoch.ConfVer, regionMeta.RegionEpoch.Version) - regionInfo := newRegionInfo(verID, intersectSpan, nil, subscribedSpan, filterLoop) + regionInfo := newRegionInfo(verID, intersectSpan, nil, subscribedSpan, task.filterLoop) + regionInfo.scanPriority = normalizeScanPriority(task.priority) // Schedule a region request to subscribe the region. - s.scheduleRegionRequest(ctx, regionInfo, taskType) + s.scheduleRegionRequest(ctx, regionInfo) nextSpan.StartKey = regionMeta.EndKey // If the nextSpan.StartKey is larger than the subscribedSpan.span.EndKey, @@ -634,13 +445,9 @@ func (s *subscriptionClient) divideSpanAndScheduleRegionRequests( } } -// scheduleRegionRequest locks the region's range and send the region to regionTaskQueue, -// which will be handled by handleRegions. -func (s *subscriptionClient) scheduleRegionRequest( - ctx context.Context, - region regionInfo, - inheritedPriority TaskType, -) { +// scheduleRegionRequest locks the region's range before submitting it to the +// request scheduler. +func (s *subscriptionClient) scheduleRegionRequest(ctx context.Context, region regionInfo) { lockRangeResult := region.subscribedSpan.rangeLock.LockRange( ctx, region.span.StartKey, region.span.EndKey, region.verID.GetID(), region.verID.GetVer()) @@ -651,29 +458,20 @@ func (s *subscriptionClient) scheduleRegionRequest( switch lockRangeResult.Status { case regionlock.LockRangeStatusSuccess: region.lockedRangeState = lockRangeResult.LockedRangeState - currentTs := s.pdClock.CurrentTS() - priority := region.subscribedSpan.priorityPolicy.resolve( - inheritedPriority, + region.scanPriority = region.subscribedSpan.priorityPolicy.resolve( + region.scanPriority, region.resolvedTs(), - oracle.GetTimeFromTS(currentTs), + s.upstream.pdClock.CurrentTime(), ) - region.scanPriority = priority.scanPriority() - s.regionTaskQueue.Push(NewRegionPriorityTask(priority, region, currentTs)) - if log.GetLevel() <= zapcore.DebugLevel { - log.Debug("cdc region scan task enqueued", - zap.Uint64("subscriptionID", uint64(region.subscribedSpan.subID)), - zap.Int64("tableID", region.subscribedSpan.span.TableID), - zap.Uint64("startTs", region.subscribedSpan.startTs), - zap.Uint64("regionID", region.verID.GetID()), - zap.Uint64("regionEpochVersion", region.verID.GetVer()), - zap.Uint64("regionEpochConfVer", region.verID.GetConfVer()), - zap.String("priority", priority.String()), - zap.String("scanPriority", region.scanPriority.String()), - zap.String("span", common.FormatTableSpan(®ion.span))) - } + s.regionScheduler.Submit(region) case regionlock.LockRangeStatusStale: for _, r := range lockRangeResult.RetryRanges { - s.scheduleRangeRequest(ctx, r, region.subscribedSpan, region.filterLoop, inheritedPriority) + s.scheduleRangeRequest(ctx, rangeTask{ + span: r, + subscribedSpan: region.subscribedSpan, + filterLoop: region.filterLoop, + priority: region.scanPriority, + }) } default: return @@ -681,19 +479,12 @@ func (s *subscriptionClient) scheduleRegionRequest( } func (s *subscriptionClient) scheduleRangeRequest( - ctx context.Context, span heartbeatpb.TableSpan, - subscribedSpan *subscribedSpan, - filterLoop bool, - inheritedPriority TaskType, + ctx context.Context, + task rangeTask, ) { select { case <-ctx.Done(): - case s.rangeTaskCh <- rangeTask{ - span: span, - subscribedSpan: subscribedSpan, - filterLoop: filterLoop, - priority: inheritedPriority, - }: + case s.rangeTaskCh <- task: } } diff --git a/logservice/logpuller/subscription_client_test.go b/logservice/logpuller/subscription_client_test.go index 18ce5cc75f..fa42a096c1 100644 --- a/logservice/logpuller/subscription_client_test.go +++ b/logservice/logpuller/subscription_client_test.go @@ -22,7 +22,6 @@ import ( "github.com/pingcap/errors" "github.com/pingcap/kvproto/pkg/cdcpb" - "github.com/pingcap/kvproto/pkg/errorpb" "github.com/pingcap/ticdc/heartbeatpb" "github.com/pingcap/ticdc/logservice/logpuller/regionlock" "github.com/pingcap/ticdc/pkg/common" @@ -58,6 +57,7 @@ func TestGenerateResolveLockTask(t *testing.T) { client := &subscriptionClient{ resolveLockTaskCh: make(chan resolveLockTask, 10), resolveLockRateLimiter: newResolveLockRateLimiter(), + memoryQuota: newMemoryQuotaController(0, 0), } client.ctx, client.cancel = context.WithCancel(context.Background()) rawSpan := heartbeatpb.TableSpan{ @@ -67,8 +67,8 @@ func TestGenerateResolveLockTask(t *testing.T) { } consumeKVEvents := func(_ []common.RawKVEntry, _ func()) bool { return false } advanceResolvedTs := func(ts uint64) {} - client.pdClock = pdutil.NewClock4Test() - client.spanRegistry = newSpanRegistry(nil, client.pdClock) + client.upstream = &upstreamHandle{pdClock: pdutil.NewClock4Test()} + client.spanRegistry = newSpanRegistry(nil, client.upstream.pdClock) span := newSubscribedSpan( client.ctx, client.resolveLockRateLimiter, @@ -80,7 +80,7 @@ func TestGenerateResolveLockTask(t *testing.T) { advanceResolvedTs, 0, false, - pdutil.NewClock4Test(), + client.upstream.pdClock, 30*time.Minute, ) client.spanRegistry.Add(span) @@ -107,16 +107,12 @@ func TestGenerateResolveLockTask(t *testing.T) { } worker := ®ionRequestWorker{ - requestCache: &requestCache{}, + tracker: newRegionTracker(), } // Lock another range, no task will be triggered before initialized. res = span.rangeLock.LockRange(context.Background(), []byte{'c'}, []byte{'d'}, 2, 100) require.Equal(t, regionlock.LockRangeStatusSuccess, res.Status) - state := newRegionFeedState(regionInfo{ - verID: tikv.NewRegionVerID(2, 1, 1), - lockedRangeState: res.LockedRangeState, - subscribedSpan: span, - }, 1, worker) + state := newRegionFeedState(regionInfo{lockedRangeState: res.LockedRangeState, subscribedSpan: span}, 1, worker, nil, nil) span.resolveStaleLocks(200) select { case <-client.resolveLockTaskCh: @@ -154,16 +150,17 @@ func TestResolveLockTaskDeduplicatedAcrossSubscribedSpans(t *testing.T) { consumeKVEvents := func(_ []common.RawKVEntry, _ func()) bool { return false } advanceResolvedTs := func(ts uint64) {} + pdClock := pdutil.NewClock4Test() span1 := newSubscribedSpan(client.ctx, client.resolveLockRateLimiter, client.resolveLockTaskCh, SubscriptionID(1), heartbeatpb.TableSpan{ TableID: 1, StartKey: []byte{'a'}, EndKey: []byte{'z'}, - }, 100, consumeKVEvents, advanceResolvedTs, 0, false, pdutil.NewClock4Test(), 30*time.Minute) + }, 100, consumeKVEvents, advanceResolvedTs, 0, false, pdClock, 30*time.Minute) span2 := newSubscribedSpan(client.ctx, client.resolveLockRateLimiter, client.resolveLockTaskCh, SubscriptionID(2), heartbeatpb.TableSpan{ TableID: 2, StartKey: []byte{'a'}, EndKey: []byte{'z'}, - }, 100, consumeKVEvents, advanceResolvedTs, 0, false, pdutil.NewClock4Test(), 30*time.Minute) + }, 100, consumeKVEvents, advanceResolvedTs, 0, false, pdClock, 30*time.Minute) res := span1.rangeLock.LockRange(context.Background(), []byte{'b'}, []byte{'c'}, 1, 100) require.Equal(t, regionlock.LockRangeStatusSuccess, res.Status) @@ -189,7 +186,7 @@ func TestResolveLockTaskDeduplicatedAcrossSubscribedSpans(t *testing.T) { } func TestHandleResolveLockTasksMetrics(t *testing.T) { - ctx, cancel := context.WithCancel(t.Context()) + ctx, cancel := context.WithCancel(context.Background()) defer cancel() resolver := &mockLockResolver{} @@ -204,10 +201,7 @@ func TestHandleResolveLockTasksMetrics(t *testing.T) { errCh <- client.handleResolveLockTasks(ctx) }() - rangeLock := regionlock.NewRangeLock(1, []byte{'a'}, []byte{'b'}, 100) - lockResult := rangeLock.LockRange(context.Background(), []byte{'a'}, []byte{'b'}, 1, 1) - require.Equal(t, regionlock.LockRangeStatusSuccess, lockResult.Status) - state := lockResult.LockedRangeState + state := ®ionlock.LockedRangeState{} state.Initialized.Store(true) state.ResolvedTs.Store(100) @@ -312,12 +306,12 @@ func TestResolveLockTaskDroppedWhenChannelFull(t *testing.T) { func TestStopTaskUsesSubscribedSpanFilterLoop(t *testing.T) { client := &subscriptionClient{ - resolveLockTaskCh: make(chan resolveLockTask, 1), - regionTaskQueue: priorityqueue.New[PriorityTask](), + resolveLockTaskCh: make(chan resolveLockTask, 1), + resolveLockRateLimiter: newResolveLockRateLimiter(), + memoryQuota: newMemoryQuotaController(0, 0), } client.ctx, client.cancel = context.WithCancel(context.Background()) defer client.cancel() - client.pdClock = pdutil.NewClock4Test() rawSpan := heartbeatpb.TableSpan{ TableID: 1, @@ -343,24 +337,26 @@ func TestStopTaskUsesSubscribedSpanFilterLoop(t *testing.T) { res := span.rangeLock.LockRange(context.Background(), rawSpan.StartKey, rawSpan.EndKey, 1, 1) require.Equal(t, regionlock.LockRangeStatusSuccess, res.Status) + const storeAddr = "store-1" + worker := ®ionRequestWorker{storeAddr: storeAddr, controlQueue: newControlQueue()} + store := ®ionRequestStore{workers: []*regionRequestWorker{worker}} + client.regionScheduler = ®ionRequestScheduler{} + client.regionScheduler.stores.Store(storeAddr, store) client.setTableStopped(span) - ctx, cancel := context.WithTimeout(context.Background(), time.Second) - defer cancel() - task, err := client.regionTaskQueue.Pop(ctx) - require.NoError(t, err) - region := task.GetRegionInfo() - require.True(t, region.isStopped()) - require.True(t, region.filterLoop) + req, ok := worker.controlQueue.tryPop() + require.True(t, ok) + require.Equal(t, SubscriptionID(1), req.subID) + require.True(t, req.filterLoop) } -func TestOnRegionFailQueuesCanceledErrorCache(t *testing.T) { +func TestRegionFailureHandlerQueuesCanceledError(t *testing.T) { client := &subscriptionClient{ - eventSink: newTestRegionEventSink(&mockDynamicStream{}), + eventSink: ®ionEventSink{ds: &mockDynamicStream{}}, } client.spanRegistry = newSpanRegistry(nil, nil) - client.failureHandler = newRegionFailureHandler(client) + client.failureHandler = newRegionFailureHandler(nil, client.onTableDrained, nil, nil) rawSpan := heartbeatpb.TableSpan{ TableID: 1, StartKey: []byte("a"), @@ -379,7 +375,7 @@ func TestOnRegionFailQueuesCanceledErrorCache(t *testing.T) { require.Equal(t, regionlock.LockRangeStatusSuccess, res2.Status) require.False(t, span.rangeLock.Stop()) - client.onRegionFail(newRegionErrorInfo(regionInfo{ + client.failureHandler.Report(newRegionErrorInfo(regionInfo{ verID: tikv.NewRegionVerID(1, 1, 1), span: heartbeatpb.TableSpan{TableID: 1, StartKey: []byte("a"), EndKey: []byte("m")}, subscribedSpan: span, @@ -389,7 +385,7 @@ func TestOnRegionFailQueuesCanceledErrorCache(t *testing.T) { require.Len(t, client.failureHandler.cache.cache, 1) require.Len(t, span.rangeLock.IterAll(nil).UnLockedRanges, 1) - client.onRegionFail(newRegionErrorInfo(regionInfo{ + client.failureHandler.Report(newRegionErrorInfo(regionInfo{ verID: tikv.NewRegionVerID(2, 1, 1), span: heartbeatpb.TableSpan{TableID: 1, StartKey: []byte("m"), EndKey: []byte("z")}, subscribedSpan: span, @@ -400,150 +396,6 @@ func TestOnRegionFailQueuesCanceledErrorCache(t *testing.T) { require.Nil(t, client.spanRegistry.Get(span.subID)) } -func TestRegionRetryScanPriority(t *testing.T) { - for _, tc := range []struct { - name string - priority cdcpb.ScanPriority - cdcErr *cdcpb.Error - everCaughtUp bool - expected TaskType - }{ - { - name: "server is busy high", - priority: cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, - cdcErr: &cdcpb.Error{ServerIsBusy: &errorpb.ServerIsBusy{}}, - expected: TaskHighPrior, - }, - { - name: "server is busy low", - priority: cdcpb.ScanPriority_SCAN_PRIORITY_LOW, - cdcErr: &cdcpb.Error{ServerIsBusy: &errorpb.ServerIsBusy{}}, - expected: TaskLowPrior, - }, - { - name: "server is busy low after catch up", - priority: cdcpb.ScanPriority_SCAN_PRIORITY_LOW, - cdcErr: &cdcpb.Error{ServerIsBusy: &errorpb.ServerIsBusy{}}, - everCaughtUp: true, - expected: TaskHighPrior, - }, - { - name: "congested high", - priority: cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, - cdcErr: &cdcpb.Error{Congested: &cdcpb.Congested{}}, - expected: TaskHighPrior, - }, - { - name: "congested low", - priority: cdcpb.ScanPriority_SCAN_PRIORITY_LOW, - cdcErr: &cdcpb.Error{Congested: &cdcpb.Congested{}}, - expected: TaskLowPrior, - }, - { - name: "unknown retry high", - priority: cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, - cdcErr: &cdcpb.Error{}, - expected: TaskHighPrior, - }, - { - name: "unknown retry low", - priority: cdcpb.ScanPriority_SCAN_PRIORITY_LOW, - cdcErr: &cdcpb.Error{}, - expected: TaskLowPrior, - }, - } { - t.Run(tc.name, func(t *testing.T) { - client := &subscriptionClient{ - regionTaskQueue: priorityqueue.New[PriorityTask](), - } - client.pdClock = pdutil.NewClock4Test() - client.pdClock.(*pdutil.Clock4Test).SetTS(oracle.GoTimeToTS(time.Now())) - client.failureHandler = newRegionFailureHandler(client) - _, span := newScanPriorityTestSpan() - span.priorityPolicy.everCaughtUp.Store(tc.everCaughtUp) - region := newScanPriorityTestRegion(span) - region.scanPriority = tc.priority - - err := client.failureHandler.handleError(context.Background(), newRegionErrorInfo(region, &eventError{err: tc.cdcErr})) - require.NoError(t, err) - - ctx, cancel := context.WithTimeout(context.Background(), time.Second) - defer cancel() - task, err := client.regionTaskQueue.Pop(ctx) - require.NoError(t, err) - require.Equal(t, tc.expected, task.(*regionPriorityTask).taskType) - require.Equal(t, tc.expected.scanPriority(), task.GetRegionInfo().scanPriority) - }) - } -} - -func TestRangeRetryPreservesScanPriority(t *testing.T) { - for _, tc := range []struct { - name string - priority cdcpb.ScanPriority - err error - expected TaskType - }{ - { - name: "epoch not match high", - priority: cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, - err: &eventError{err: &cdcpb.Error{EpochNotMatch: &errorpb.EpochNotMatch{}}}, - expected: TaskHighPrior, - }, - { - name: "epoch not match low", - priority: cdcpb.ScanPriority_SCAN_PRIORITY_LOW, - err: &eventError{err: &cdcpb.Error{EpochNotMatch: &errorpb.EpochNotMatch{}}}, - expected: TaskLowPrior, - }, - { - name: "region not found high", - priority: cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, - err: &eventError{err: &cdcpb.Error{RegionNotFound: &errorpb.RegionNotFound{}}}, - expected: TaskHighPrior, - }, - { - name: "region not found low", - priority: cdcpb.ScanPriority_SCAN_PRIORITY_LOW, - err: &eventError{err: &cdcpb.Error{RegionNotFound: &errorpb.RegionNotFound{}}}, - expected: TaskLowPrior, - }, - { - name: "rpc context unavailable high", - priority: cdcpb.ScanPriority_SCAN_PRIORITY_HIGH, - err: &rpcCtxUnavailableErr{verID: tikv.NewRegionVerID(1, 1, 1)}, - expected: TaskHighPrior, - }, - { - name: "rpc context unavailable low", - priority: cdcpb.ScanPriority_SCAN_PRIORITY_LOW, - err: &rpcCtxUnavailableErr{verID: tikv.NewRegionVerID(1, 1, 1)}, - expected: TaskLowPrior, - }, - } { - t.Run(tc.name, func(t *testing.T) { - client := &subscriptionClient{ - rangeTaskCh: make(chan rangeTask, 1), - } - client.failureHandler = newRegionFailureHandler(client) - rawSpan, span := newScanPriorityTestSpan() - region := newScanPriorityTestRegion(span) - region.scanPriority = tc.priority - - err := client.failureHandler.handleError(context.Background(), newRegionErrorInfo(region, tc.err)) - require.NoError(t, err) - - select { - case task := <-client.rangeTaskCh: - require.Equal(t, tc.expected, task.priority) - require.Equal(t, rawSpan, task.span) - case <-time.After(time.Second): - require.Fail(t, "expected range retry task") - } - }) - } -} - type mockDynamicStream struct{} func (s *mockDynamicStream) Start() {} @@ -574,48 +426,49 @@ func (s *mockDynamicStream) GetMetrics() dynstream.Metrics[int, SubscriptionID] return dynstream.Metrics[int, SubscriptionID]{} } -func newScanPriorityTestSpan() (heartbeatpb.TableSpan, *subscribedSpan) { - rawSpan := heartbeatpb.TableSpan{ - TableID: 1, - StartKey: []byte("a"), - EndKey: []byte("z"), - } - span := &subscribedSpan{ - subID: SubscriptionID(1), - span: rawSpan, - rangeLock: regionlock.NewRangeLock(1, rawSpan.StartKey, rawSpan.EndKey, 100), - priorityPolicy: newTestScanPriorityPolicy(), - } - return rawSpan, span -} - -func newScanPriorityTestRegion(span *subscribedSpan) regionInfo { - return newRegionInfo(tikv.NewRegionVerID(1, 1, 1), span.span, nil, span, false) -} - -func newTestScanPriorityPolicy() scanPriorityPolicy { - return newScanPriorityPolicy(pdutil.NewClock4Test(), 30*time.Minute) -} +func TestRegionEventSinkPushUnblocksOnClientClose(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() -func TestPushRegionEventToDSUnblocksOnClose(t *testing.T) { - sink := newTestRegionEventSink(&mockDynamicStream{}) - client := &subscriptionClient{ - eventSink: sink, - regionTaskQueue: priorityqueue.New[PriorityTask](), + quota := newMemoryQuotaController(10, 8) + span := &subscribedSpan{subID: 1} + require.True(t, quota.AcquireEvent(ctx, span, 20)) + t.Cleanup(func() { quota.ReleaseEvent(20) }) + + sink := ®ionEventSink{ + ctx: ctx, + ds: &mockDynamicStream{}, + memoryQuota: quota, + } + client := &subscriptionClient{eventSink: sink} + client.regionScheduler = ®ionRequestScheduler{ + taskQueue: priorityqueue.New[*regionPriorityTask](), + } + client.ctx = ctx + client.cancel = cancel + + event := regionEvent{ + states: []*regionFeedState{{ + region: regionInfo{subscribedSpan: span}, + }}, + entries: &cdcpb.Event_Entries_{ + Entries: &cdcpb.Event_Entries{ + Entries: []*cdcpb.Event_Row{{ + Key: []byte("key"), + Value: []byte("value"), + }}, + }, + }, } - client.ctx, client.cancel = context.WithCancel(context.Background()) - - sink.paused.Store(true) - done := make(chan struct{}) go func() { - client.pushRegionEventToDS(SubscriptionID(1), regionEvent{}) + sink.Push(SubscriptionID(1), event) close(done) }() select { case <-done: - t.Fatal("pushRegionEventToDS should block when paused") + t.Fatal("pushRegionEventToDS should block when event memory is exhausted") case <-time.After(100 * time.Millisecond): } @@ -628,41 +481,6 @@ func TestPushRegionEventToDSUnblocksOnClose(t *testing.T) { } } -func TestEnqueueRegionToAllStoresRetryWhenCacheFull(t *testing.T) { - ctx := context.Background() - client := &subscriptionClient{} - - worker := ®ionRequestWorker{ - requestCache: newRequestCache(1), - } - store := &requestedStore{storeAddr: "store-1"} - store.requestWorkers.s = []*regionRequestWorker{worker} - client.stores.Store(store.storeAddr, store) - - dummyRegion := regionInfo{ - subscribedSpan: &subscribedSpan{subID: SubscriptionID(2)}, - lockedRangeState: ®ionlock.LockedRangeState{}, - } - ok, err := worker.add(ctx, dummyRegion, true) - require.NoError(t, err) - require.True(t, ok) - - stopRegion := regionInfo{ - subscribedSpan: &subscribedSpan{subID: SubscriptionID(1)}, - } - enqueued, err := client.enqueueRegionToAllStores(ctx, stopRegion) - require.NoError(t, err) - require.False(t, enqueued) - - <-worker.requestCache.pendingQueue - worker.requestCache.markDone() - - enqueued, err = client.enqueueRegionToAllStores(ctx, stopRegion) - require.NoError(t, err) - require.True(t, enqueued) - require.Equal(t, 1, len(worker.requestCache.pendingQueue)) -} - func TestSubscriptionWithFailedTiKV(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) mockPDClock := pdutil.NewClock4Test() @@ -692,11 +510,7 @@ func TestSubscriptionWithFailedTiKV(t *testing.T) { // bootstrap cluster with a region which leader is in invalid store. cluster.Bootstrap(11, []uint64{1, 2, 3}, []uint64{4, 5, 6}, 6) - clientConfig := &SubscriptionClientConfig{ - RegionRequestWorkerPerStore: 2, - } client := NewSubscriptionClient( - clientConfig, pdClient, nil, // we don't need it in this unittest, so we can pass nil &security.Credential{}, diff --git a/metrics/grafana/ticdc_new_arch.json b/metrics/grafana/ticdc_new_arch.json index 0a3f4bbd37..c92a6c9381 100644 --- a/metrics/grafana/ticdc_new_arch.json +++ b/metrics/grafana/ticdc_new_arch.json @@ -8119,7 +8119,7 @@ "dashLength": 10, "dashes": false, "datasource": "${DS_TEST-CLUSTER}", - "description": "Log puller memory quota", + "description": "Actual event memory, estimated in-flight scan memory, and the configured memory limits.", "fieldConfig": { "defaults": {}, "overrides": [] @@ -8162,7 +8162,7 @@ "targets": [ { "exemplar": true, - "expr": "sum(ticdc_dynamic_stream_memory_usage{k8s_cluster=~\"$k8s_cluster\", tidb_cluster=\"$tidb_cluster\", instance=~\"$ticdc_instance\", module=~\"log-puller\"}) by (instance, type)", + "expr": "sum(ticdc_log_puller_memory_quota{k8s_cluster=~\"$k8s_cluster\", tidb_cluster=\"$tidb_cluster\", instance=~\"$ticdc_instance\"}) by (instance, type)", "interval": "", "legendFormat": "{{instance}}-{{type}}", "refId": "A" @@ -8172,7 +8172,7 @@ "timeFrom": null, "timeRegions": [], "timeShift": null, - "title": "Memory Quota", + "title": "Memory Quota Usage", "tooltip": { "shared": true, "sort": 0, @@ -8815,6 +8815,214 @@ "align": false, "alignLevel": null } + }, + { + "aliasColors": {}, + "bars": false, + "dashLength": 10, + "dashes": false, + "datasource": "${DS_TEST-CLUSTER}", + "description": "Event receivers blocked at the memory hard limit and region scans blocked at the scan admission gate.", + "fieldConfig": { + "defaults": {}, + "overrides": [] + }, + "fill": 0, + "fillGradient": 0, + "gridPos": { + "h": 8, + "w": 12, + "x": 0, + "y": 57 + }, + "hiddenSeries": false, + "id": 26001, + "legend": { + "alignAsTable": true, + "avg": false, + "current": true, + "max": true, + "min": false, + "show": true, + "total": false, + "values": true + }, + "lines": true, + "linewidth": 1, + "nullPointMode": "null", + "options": { + "alertThreshold": true + }, + "percentage": false, + "pluginVersion": "7.5.17", + "pointradius": 2, + "points": false, + "renderer": "flot", + "seriesOverrides": [], + "spaceLength": 10, + "stack": false, + "steppedLine": false, + "targets": [ + { + "exemplar": true, + "expr": "ticdc_log_puller_memory_quota_event_waiter_count{k8s_cluster=~\"$k8s_cluster\", tidb_cluster=\"$tidb_cluster\", instance=~\"$ticdc_instance\"}", + "interval": "", + "legendFormat": "{{instance}}-event-waiters", + "refId": "A" + }, + { + "exemplar": true, + "expr": "ticdc_log_puller_memory_quota_scan_waiter_count{k8s_cluster=~\"$k8s_cluster\", tidb_cluster=\"$tidb_cluster\", instance=~\"$ticdc_instance\"}", + "interval": "", + "legendFormat": "{{instance}}-scan-waiters", + "refId": "B" + } + ], + "thresholds": [], + "timeFrom": null, + "timeRegions": [], + "timeShift": null, + "title": "Memory Quota Waiters", + "tooltip": { + "shared": true, + "sort": 0, + "value_type": "individual" + }, + "type": "graph", + "xaxis": { + "buckets": null, + "mode": "time", + "name": null, + "show": true, + "values": [] + }, + "yaxes": [ + { + "format": "short", + "logBase": 1, + "min": "0", + "show": true + }, + { + "format": "short", + "logBase": 1, + "show": false + } + ], + "yaxis": { + "align": false + } + }, + { + "aliasColors": {}, + "bars": false, + "dashLength": 10, + "dashes": false, + "datasource": "${DS_TEST-CLUSTER}", + "description": "Time spent waiting at the event hard limit or the scan admission gate.", + "fieldConfig": { + "defaults": {}, + "overrides": [] + }, + "fill": 0, + "fillGradient": 0, + "gridPos": { + "h": 8, + "w": 12, + "x": 12, + "y": 57 + }, + "hiddenSeries": false, + "id": 26002, + "legend": { + "alignAsTable": true, + "avg": false, + "current": true, + "max": true, + "min": false, + "show": true, + "total": false, + "values": true + }, + "lines": true, + "linewidth": 1, + "nullPointMode": "null", + "options": { + "alertThreshold": true + }, + "percentage": false, + "pluginVersion": "7.5.17", + "pointradius": 2, + "points": false, + "renderer": "flot", + "seriesOverrides": [], + "spaceLength": 10, + "stack": false, + "steppedLine": false, + "targets": [ + { + "exemplar": true, + "expr": "histogram_quantile(0.99, sum(rate(ticdc_log_puller_memory_quota_event_wait_duration_bucket{k8s_cluster=~\"$k8s_cluster\", tidb_cluster=\"$tidb_cluster\", instance=~\"$ticdc_instance\"}[1m])) by (le, instance))", + "interval": "", + "legendFormat": "{{instance}}-event-p99", + "refId": "A" + }, + { + "exemplar": true, + "expr": "sum(rate(ticdc_log_puller_memory_quota_event_wait_duration_sum{k8s_cluster=~\"$k8s_cluster\", tidb_cluster=\"$tidb_cluster\", instance=~\"$ticdc_instance\"}[1m])) by (instance) / sum(rate(ticdc_log_puller_memory_quota_event_wait_duration_count{k8s_cluster=~\"$k8s_cluster\", tidb_cluster=\"$tidb_cluster\", instance=~\"$ticdc_instance\"}[1m])) by (instance)", + "interval": "", + "legendFormat": "{{instance}}-event-avg", + "refId": "B" + }, + { + "exemplar": true, + "expr": "histogram_quantile(0.99, sum(rate(ticdc_log_puller_memory_quota_scan_wait_duration_bucket{k8s_cluster=~\"$k8s_cluster\", tidb_cluster=\"$tidb_cluster\", instance=~\"$ticdc_instance\"}[1m])) by (le, instance))", + "interval": "", + "legendFormat": "{{instance}}-scan-p99", + "refId": "C" + }, + { + "exemplar": true, + "expr": "sum(rate(ticdc_log_puller_memory_quota_scan_wait_duration_sum{k8s_cluster=~\"$k8s_cluster\", tidb_cluster=\"$tidb_cluster\", instance=~\"$ticdc_instance\"}[1m])) by (instance) / sum(rate(ticdc_log_puller_memory_quota_scan_wait_duration_count{k8s_cluster=~\"$k8s_cluster\", tidb_cluster=\"$tidb_cluster\", instance=~\"$ticdc_instance\"}[1m])) by (instance)", + "interval": "", + "legendFormat": "{{instance}}-scan-avg", + "refId": "D" + } + ], + "thresholds": [], + "timeFrom": null, + "timeRegions": [], + "timeShift": null, + "title": "Memory Quota Wait Duration", + "tooltip": { + "shared": true, + "sort": 0, + "value_type": "individual" + }, + "type": "graph", + "xaxis": { + "buckets": null, + "mode": "time", + "name": null, + "show": true, + "values": [] + }, + "yaxes": [ + { + "format": "s", + "logBase": 1, + "min": "0", + "show": true + }, + { + "format": "short", + "logBase": 1, + "show": false + } + ], + "yaxis": { + "align": false + } } ], "title": "Log Puller", diff --git a/metrics/nextgengrafana/ticdc_new_arch_next_gen.json b/metrics/nextgengrafana/ticdc_new_arch_next_gen.json index f920ac2ccd..b684f1d915 100644 --- a/metrics/nextgengrafana/ticdc_new_arch_next_gen.json +++ b/metrics/nextgengrafana/ticdc_new_arch_next_gen.json @@ -8119,7 +8119,7 @@ "dashLength": 10, "dashes": false, "datasource": "${DS_TEST-CLUSTER}", - "description": "Log puller memory quota", + "description": "Actual event memory, estimated in-flight scan memory, and the configured memory limits.", "fieldConfig": { "defaults": {}, "overrides": [] @@ -8162,7 +8162,7 @@ "targets": [ { "exemplar": true, - "expr": "sum(ticdc_dynamic_stream_memory_usage{k8s_cluster=~\"$k8s_cluster\", sharedpool_id=\"$tidb_cluster\", instance=~\"$ticdc_instance\", module=~\"log-puller\"}) by (instance, type)", + "expr": "sum(ticdc_log_puller_memory_quota{k8s_cluster=~\"$k8s_cluster\", sharedpool_id=\"$tidb_cluster\", instance=~\"$ticdc_instance\"}) by (instance, type)", "interval": "", "legendFormat": "{{instance}}-{{type}}", "refId": "A" @@ -8172,7 +8172,7 @@ "timeFrom": null, "timeRegions": [], "timeShift": null, - "title": "Memory Quota", + "title": "Memory Quota Usage", "tooltip": { "shared": true, "sort": 0, @@ -8815,6 +8815,214 @@ "align": false, "alignLevel": null } + }, + { + "aliasColors": {}, + "bars": false, + "dashLength": 10, + "dashes": false, + "datasource": "${DS_TEST-CLUSTER}", + "description": "Event receivers blocked at the memory hard limit and region scans blocked at the scan admission gate.", + "fieldConfig": { + "defaults": {}, + "overrides": [] + }, + "fill": 0, + "fillGradient": 0, + "gridPos": { + "h": 8, + "w": 12, + "x": 0, + "y": 57 + }, + "hiddenSeries": false, + "id": 26001, + "legend": { + "alignAsTable": true, + "avg": false, + "current": true, + "max": true, + "min": false, + "show": true, + "total": false, + "values": true + }, + "lines": true, + "linewidth": 1, + "nullPointMode": "null", + "options": { + "alertThreshold": true + }, + "percentage": false, + "pluginVersion": "7.5.17", + "pointradius": 2, + "points": false, + "renderer": "flot", + "seriesOverrides": [], + "spaceLength": 10, + "stack": false, + "steppedLine": false, + "targets": [ + { + "exemplar": true, + "expr": "ticdc_log_puller_memory_quota_event_waiter_count{k8s_cluster=~\"$k8s_cluster\", sharedpool_id=\"$tidb_cluster\", instance=~\"$ticdc_instance\"}", + "interval": "", + "legendFormat": "{{instance}}-event-waiters", + "refId": "A" + }, + { + "exemplar": true, + "expr": "ticdc_log_puller_memory_quota_scan_waiter_count{k8s_cluster=~\"$k8s_cluster\", sharedpool_id=\"$tidb_cluster\", instance=~\"$ticdc_instance\"}", + "interval": "", + "legendFormat": "{{instance}}-scan-waiters", + "refId": "B" + } + ], + "thresholds": [], + "timeFrom": null, + "timeRegions": [], + "timeShift": null, + "title": "Memory Quota Waiters", + "tooltip": { + "shared": true, + "sort": 0, + "value_type": "individual" + }, + "type": "graph", + "xaxis": { + "buckets": null, + "mode": "time", + "name": null, + "show": true, + "values": [] + }, + "yaxes": [ + { + "format": "short", + "logBase": 1, + "min": "0", + "show": true + }, + { + "format": "short", + "logBase": 1, + "show": false + } + ], + "yaxis": { + "align": false + } + }, + { + "aliasColors": {}, + "bars": false, + "dashLength": 10, + "dashes": false, + "datasource": "${DS_TEST-CLUSTER}", + "description": "Time spent waiting at the event hard limit or the scan admission gate.", + "fieldConfig": { + "defaults": {}, + "overrides": [] + }, + "fill": 0, + "fillGradient": 0, + "gridPos": { + "h": 8, + "w": 12, + "x": 12, + "y": 57 + }, + "hiddenSeries": false, + "id": 26002, + "legend": { + "alignAsTable": true, + "avg": false, + "current": true, + "max": true, + "min": false, + "show": true, + "total": false, + "values": true + }, + "lines": true, + "linewidth": 1, + "nullPointMode": "null", + "options": { + "alertThreshold": true + }, + "percentage": false, + "pluginVersion": "7.5.17", + "pointradius": 2, + "points": false, + "renderer": "flot", + "seriesOverrides": [], + "spaceLength": 10, + "stack": false, + "steppedLine": false, + "targets": [ + { + "exemplar": true, + "expr": "histogram_quantile(0.99, sum(rate(ticdc_log_puller_memory_quota_event_wait_duration_bucket{k8s_cluster=~\"$k8s_cluster\", sharedpool_id=\"$tidb_cluster\", instance=~\"$ticdc_instance\"}[1m])) by (le, instance))", + "interval": "", + "legendFormat": "{{instance}}-event-p99", + "refId": "A" + }, + { + "exemplar": true, + "expr": "sum(rate(ticdc_log_puller_memory_quota_event_wait_duration_sum{k8s_cluster=~\"$k8s_cluster\", sharedpool_id=\"$tidb_cluster\", instance=~\"$ticdc_instance\"}[1m])) by (instance) / sum(rate(ticdc_log_puller_memory_quota_event_wait_duration_count{k8s_cluster=~\"$k8s_cluster\", sharedpool_id=\"$tidb_cluster\", instance=~\"$ticdc_instance\"}[1m])) by (instance)", + "interval": "", + "legendFormat": "{{instance}}-event-avg", + "refId": "B" + }, + { + "exemplar": true, + "expr": "histogram_quantile(0.99, sum(rate(ticdc_log_puller_memory_quota_scan_wait_duration_bucket{k8s_cluster=~\"$k8s_cluster\", sharedpool_id=\"$tidb_cluster\", instance=~\"$ticdc_instance\"}[1m])) by (le, instance))", + "interval": "", + "legendFormat": "{{instance}}-scan-p99", + "refId": "C" + }, + { + "exemplar": true, + "expr": "sum(rate(ticdc_log_puller_memory_quota_scan_wait_duration_sum{k8s_cluster=~\"$k8s_cluster\", sharedpool_id=\"$tidb_cluster\", instance=~\"$ticdc_instance\"}[1m])) by (instance) / sum(rate(ticdc_log_puller_memory_quota_scan_wait_duration_count{k8s_cluster=~\"$k8s_cluster\", sharedpool_id=\"$tidb_cluster\", instance=~\"$ticdc_instance\"}[1m])) by (instance)", + "interval": "", + "legendFormat": "{{instance}}-scan-avg", + "refId": "D" + } + ], + "thresholds": [], + "timeFrom": null, + "timeRegions": [], + "timeShift": null, + "title": "Memory Quota Wait Duration", + "tooltip": { + "shared": true, + "sort": 0, + "value_type": "individual" + }, + "type": "graph", + "xaxis": { + "buckets": null, + "mode": "time", + "name": null, + "show": true, + "values": [] + }, + "yaxes": [ + { + "format": "s", + "logBase": 1, + "min": "0", + "show": true + }, + { + "format": "short", + "logBase": 1, + "show": false + } + ], + "yaxis": { + "align": false + } } ], "title": "Log Puller", diff --git a/pkg/config/debug.go b/pkg/config/debug.go index 9186170b52..eb7e45c472 100644 --- a/pkg/config/debug.go +++ b/pkg/config/debug.go @@ -17,12 +17,16 @@ import ( "time" "github.com/pingcap/errors" + "github.com/pingcap/log" + "go.uber.org/zap" ) const ( // DefaultOldStartTsScanLowPriorityThreshold is the default lag threshold for // classifying scan tasks as low priority. - DefaultOldStartTsScanLowPriorityThreshold = 30 * time.Minute + DefaultOldStartTsScanLowPriorityThreshold = 10 * time.Minute + defaultLogPullerMemoryQuota uint64 = 1024 * 1024 * 1024 + defaultLogPullerScanBaseSize uint64 = 8 * 1024 * 1024 ) // DebugConfig represents config for ticdc unexposed feature configurations @@ -69,34 +73,68 @@ type PullerConfig struct { // LogRegionDetails determines whether logs Region details or not in puller and kv-client. LogRegionDetails bool `toml:"log-region-details" json:"log_region_details"` - // PendingRegionRequestQueueSize is the total size of the pending region request queue shared across - // all puller workers connecting to a single TiKV store. This size is divided equally among all workers. - // For example, if PendingRegionRequestQueueSize is 32 and there are 8 workers connecting to the same store, - // each worker's queue size will be 32 / 8 = 4. + // PendingRegionRequestQueueSize is the approximate normal initial-scan window + // for one TiKV store. It is divided among the store's puller workers. PendingRegionRequestQueueSize int `toml:"pending-region-request-queue-size" json:"pending_region_request_queue_size"` + // RegionRequestMaxWindowMultiplier controls the maximum window available to + // previously initialized regions and regions whose scan lag is below the + // configured low-priority threshold. + // The approximate maximum store window is PendingRegionRequestQueueSize + // multiplied by this value. + RegionRequestMaxWindowMultiplier int `toml:"region-request-max-window-multiplier" json:"region_request_max_window_multiplier"` // OldStartTsScanLowPriorityThreshold is the lag threshold for scan priority. // Scans within this threshold are scheduled as high priority. Older scans // remain low priority until their span catches up once. OldStartTsScanLowPriorityThreshold TomlDuration `toml:"old-start-ts-scan-low-priority-threshold" json:"old_start_ts_scan_low_priority_threshold"` + // MemoryQuota is the log puller's local soft memory limit in bytes. + MemoryQuota uint64 `toml:"memory-quota" json:"memory_quota"` + // ScanBaseSize is the base memory estimate for one admitted initial scan. + ScanBaseSize uint64 `toml:"scan-base-size" json:"scan_base_size"` } // NewDefaultPullerConfig return the default puller configuration func NewDefaultPullerConfig() *PullerConfig { return &PullerConfig{ - EnableResolvedTsStuckDetection: false, - ResolvedTsStuckInterval: TomlDuration(5 * time.Minute), - LogRegionDetails: false, - PendingRegionRequestQueueSize: 32, // This value is chosen to reduce the impact of new changefeeds on existing ones. + EnableResolvedTsStuckDetection: false, + ResolvedTsStuckInterval: TomlDuration(5 * time.Minute), + LogRegionDetails: false, + PendingRegionRequestQueueSize: 32, // This value is chosen to reduce the impact of new changefeeds on existing ones. + RegionRequestMaxWindowMultiplier: 4, OldStartTsScanLowPriorityThreshold: TomlDuration( DefaultOldStartTsScanLowPriorityThreshold), + MemoryQuota: defaultLogPullerMemoryQuota, + ScanBaseSize: defaultLogPullerScanBaseSize, } } // ValidateAndAdjust validates and adjusts puller configuration. func (c *PullerConfig) ValidateAndAdjust() { + defaultCfg := NewDefaultPullerConfig() + if c.PendingRegionRequestQueueSize <= 0 { + log.Warn("pending region request queue size must be positive, use default value", + zap.Int("value", c.PendingRegionRequestQueueSize), + zap.Int("default", defaultCfg.PendingRegionRequestQueueSize)) + c.PendingRegionRequestQueueSize = defaultCfg.PendingRegionRequestQueueSize + } + if c.RegionRequestMaxWindowMultiplier <= 0 { + log.Warn("region request max window multiplier must be positive, use default value", + zap.Int("value", c.RegionRequestMaxWindowMultiplier), + zap.Int("default", defaultCfg.RegionRequestMaxWindowMultiplier)) + c.RegionRequestMaxWindowMultiplier = defaultCfg.RegionRequestMaxWindowMultiplier + } if c.OldStartTsScanLowPriorityThreshold <= 0 { c.OldStartTsScanLowPriorityThreshold = TomlDuration(DefaultOldStartTsScanLowPriorityThreshold) } + if c.MemoryQuota == 0 { + log.Warn("log puller memory quota must be positive, use default value", + zap.Uint64("default", defaultCfg.MemoryQuota)) + c.MemoryQuota = defaultCfg.MemoryQuota + } + if c.ScanBaseSize == 0 { + log.Warn("log puller scan base size must be positive, use default value", + zap.Uint64("default", defaultCfg.ScanBaseSize)) + c.ScanBaseSize = defaultCfg.ScanBaseSize + } } type EventStoreConfig struct { diff --git a/pkg/config/debug_test.go b/pkg/config/debug_test.go new file mode 100644 index 0000000000..82904dcabf --- /dev/null +++ b/pkg/config/debug_test.go @@ -0,0 +1,39 @@ +// Copyright 2026 PingCAP, Inc. +// +// 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 config + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestPullerConfigValidateAndAdjustRegionRequestWindow(t *testing.T) { + defaultCfg := NewDefaultPullerConfig() + require.Equal(t, 32, defaultCfg.PendingRegionRequestQueueSize) + require.Equal(t, 4, defaultCfg.RegionRequestMaxWindowMultiplier) + require.Equal(t, uint64(1024*1024*1024), defaultCfg.MemoryQuota) + require.Equal(t, uint64(8*1024*1024), defaultCfg.ScanBaseSize) + + cfg := &PullerConfig{ + PendingRegionRequestQueueSize: -1, + RegionRequestMaxWindowMultiplier: 0, + } + cfg.ValidateAndAdjust() + require.Equal(t, defaultCfg.PendingRegionRequestQueueSize, cfg.PendingRegionRequestQueueSize) + require.Equal(t, defaultCfg.RegionRequestMaxWindowMultiplier, cfg.RegionRequestMaxWindowMultiplier) + require.Equal(t, defaultCfg.MemoryQuota, cfg.MemoryQuota) + require.Equal(t, defaultCfg.ScanBaseSize, cfg.ScanBaseSize) +} diff --git a/pkg/metrics/log_puller.go b/pkg/metrics/log_puller.go index e874aa4428..08f3ecb137 100644 --- a/pkg/metrics/log_puller.go +++ b/pkg/metrics/log_puller.go @@ -64,6 +64,43 @@ var ( Name: "resolved_ts_lag", Help: "The lag of resolved ts", }) + LogPullerMemoryQuota = prometheus.NewGaugeVec( + prometheus.GaugeOpts{ + Namespace: "ticdc", + Subsystem: "log_puller", + Name: "memory_quota", + Help: "The log puller local memory quota usage.", + }, []string{"type"}) + LogPullerMemoryQuotaEventWaiterCount = prometheus.NewGauge( + prometheus.GaugeOpts{ + Namespace: "ticdc", + Subsystem: "log_puller", + Name: "memory_quota_event_waiter_count", + Help: "The number of event receivers waiting at the log puller memory hard limit.", + }) + LogPullerMemoryQuotaEventWaitDuration = prometheus.NewHistogram( + prometheus.HistogramOpts{ + Namespace: "ticdc", + Subsystem: "log_puller", + Name: "memory_quota_event_wait_duration", + Help: "The duration in seconds that an event receiver waits at the log puller memory hard limit.", + Buckets: prometheus.ExponentialBuckets(0.001, 2, 24), + }) + LogPullerMemoryQuotaScanWaiterCount = prometheus.NewGauge( + prometheus.GaugeOpts{ + Namespace: "ticdc", + Subsystem: "log_puller", + Name: "memory_quota_scan_waiter_count", + Help: "The number of region scans waiting at the log puller memory quota gate.", + }) + LogPullerMemoryQuotaScanWaitDuration = prometheus.NewHistogram( + prometheus.HistogramOpts{ + Namespace: "ticdc", + Subsystem: "log_puller", + Name: "memory_quota_scan_wait_duration", + Help: "The duration in seconds that a region scan waits at the log puller memory quota gate.", + Buckets: prometheus.ExponentialBuckets(0.001, 2, 24), + }) SubscriptionClientResolvedTsLagGauge = prometheus.NewGauge( prometheus.GaugeOpts{ @@ -78,7 +115,7 @@ var ( Namespace: "ticdc", Subsystem: "subscription_client", Name: "requested_region_count", - Help: "The number of requested regions", + Help: "The number of region requests by state.", }, []string{"state"}) RegionRequestFinishScanDuration = prometheus.NewHistogram( prometheus.HistogramOpts{ @@ -164,6 +201,11 @@ func initLogPullerMetrics(registry *prometheus.Registry) { registry.MustRegister(LogPullerPrewriteCacheRowNum) registry.MustRegister(LogPullerMatcherCount) registry.MustRegister(LogPullerResolvedTsLag) + registry.MustRegister(LogPullerMemoryQuota) + registry.MustRegister(LogPullerMemoryQuotaEventWaiterCount) + registry.MustRegister(LogPullerMemoryQuotaEventWaitDuration) + registry.MustRegister(LogPullerMemoryQuotaScanWaiterCount) + registry.MustRegister(LogPullerMemoryQuotaScanWaitDuration) registry.MustRegister(SubscriptionClientRequestedRegionCount) registry.MustRegister(SubscriptionClientAddRegionRequestDuration) registry.MustRegister(RegionRequestFinishScanDuration) diff --git a/server/server.go b/server/server.go index b9dc5edad8..8a604d35cc 100644 --- a/server/server.go +++ b/server/server.go @@ -198,9 +198,7 @@ func (c *server) initialize(ctx context.Context) error { conf := config.GetGlobalServerConfig() schemaStore := schemastore.New(conf.DataDir, c.pdClient) subscriptionClient := logpuller.NewSubscriptionClient( - &logpuller.SubscriptionClientConfig{ - RegionRequestWorkerPerStore: 8, - }, c.pdClient, + c.pdClient, txnutil.NewLockerResolver(), c.security, ) diff --git a/utils/notifyqueue/notify_queue.go b/utils/notifyqueue/notify_queue.go new file mode 100644 index 0000000000..9b2fea8d70 --- /dev/null +++ b/utils/notifyqueue/notify_queue.go @@ -0,0 +1,78 @@ +// Copyright 2026 PingCAP, Inc. +// +// 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 notifyqueue provides a FIFO queue with a selectable ready signal. +package notifyqueue + +import "github.com/pingcap/ticdc/utils/deque" + +// Queue is a FIFO queue with a wake-up channel. +// +// Queue is not safe for concurrent use. Callers must serialize Push, TryPop, +// Drain, and Len by themselves. +// +// Ready returns a wake-up hint. Receiving from it does not guarantee TryPop +// will return an item, and callers must re-check the queue state after wake-up. +type Queue[T any] struct { + queue *deque.Deque[T] + ready chan struct{} +} + +// New creates an empty Queue. +func New[T any]() *Queue[T] { + return &Queue[T]{ + queue: deque.NewDequeDefault[T](), + ready: make(chan struct{}, 1), + } +} + +// Push appends an item and signals Ready. +func (q *Queue[T]) Push(item T) { + q.queue.PushBack(item) + q.signal() +} + +// TryPop pops one item from the front of the queue. +func (q *Queue[T]) TryPop() (T, bool) { + return q.queue.PopFront() +} + +// Drain removes and returns all queued items in FIFO order. +func (q *Queue[T]) Drain() []T { + items := make([]T, 0, q.queue.Length()) + for { + item, ok := q.queue.PopFront() + if !ok { + return items + } + items = append(items, item) + } +} + +// Len returns the number of queued items. +func (q *Queue[T]) Len() int { + return q.queue.Length() +} + +// Ready returns a wake-up channel for consumers. +func (q *Queue[T]) Ready() <-chan struct{} { + return q.ready +} + +func (q *Queue[T]) signal() { + select { + case q.ready <- struct{}{}: + default: + } +} diff --git a/utils/notifyqueue/notify_queue_test.go b/utils/notifyqueue/notify_queue_test.go new file mode 100644 index 0000000000..228e9e563f --- /dev/null +++ b/utils/notifyqueue/notify_queue_test.go @@ -0,0 +1,84 @@ +// Copyright 2026 PingCAP, Inc. +// +// 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 notifyqueue + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestQueuePushPopAndReady(t *testing.T) { + q := New[int]() + + select { + case <-q.Ready(): + t.Fatal("empty queue should not be ready") + default: + } + + q.Push(1) + q.Push(2) + require.Equal(t, 2, q.Len()) + + select { + case <-q.Ready(): + default: + t.Fatal("push should signal ready") + } + + v, ok := q.TryPop() + require.True(t, ok) + require.Equal(t, 1, v) + v, ok = q.TryPop() + require.True(t, ok) + require.Equal(t, 2, v) + require.Equal(t, 0, q.Len()) + + _, ok = q.TryPop() + require.False(t, ok) +} + +func TestQueueReadyIsCoalesced(t *testing.T) { + q := New[int]() + + q.Push(1) + q.Push(2) + + select { + case <-q.Ready(): + default: + t.Fatal("push should signal ready") + } + + select { + case <-q.Ready(): + t.Fatal("ready signal should be coalesced") + default: + } +} + +func TestQueueDrain(t *testing.T) { + q := New[int]() + q.Push(1) + q.Push(2) + q.Push(3) + + require.Equal(t, []int{1, 2, 3}, q.Drain()) + require.Equal(t, 0, q.Len()) + + _, ok := q.TryPop() + require.False(t, ok) +} diff --git a/utils/priorityqueue/priority_queue_test.go b/utils/priorityqueue/priority_queue_test.go index 7e9dba6d1e..b501c456d2 100644 --- a/utils/priorityqueue/priority_queue_test.go +++ b/utils/priorityqueue/priority_queue_test.go @@ -83,22 +83,42 @@ func TestQueuePopBlocking(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) defer cancel() - start := time.Now() - task, err := q.Pop(ctx) - require.ErrorIs(t, err, context.DeadlineExceeded) - require.Nil(t, task) - require.GreaterOrEqual(t, time.Since(start), 50*time.Millisecond) + type popResult struct { + task *mockItem + err error + } + resultCh := make(chan popResult, 1) + go func() { + task, err := q.Pop(ctx) + resultCh <- popResult{task: task, err: err} + }() + select { + case result := <-resultCh: + require.ErrorIs(t, result.err, context.DeadlineExceeded) + require.Nil(t, result.task) + case <-time.After(time.Second): + t.Fatal("Pop did not return after context timeout") + } + + resultCh = make(chan popResult, 1) go func() { time.Sleep(50 * time.Millisecond) q.Push(newMockItem(10, "task1")) }() - start = time.Now() - task, err = q.Pop(context.Background()) - require.NoError(t, err) - require.Equal(t, "task1", task.description) - require.GreaterOrEqual(t, time.Since(start), 50*time.Millisecond) + go func() { + task, err := q.Pop(context.Background()) + resultCh <- popResult{task: task, err: err} + }() + + select { + case result := <-resultCh: + require.NoError(t, result.err) + require.Equal(t, "task1", result.task.description) + case <-time.After(time.Second): + t.Fatal("Pop did not return after Push") + } } func TestQueueTryPopAndUpdateExistingItem(t *testing.T) {