From 51353fac93c227fb8307afe2c8725164a3045ba6 Mon Sep 17 00:00:00 2001 From: aptend Date: Mon, 20 Jul 2026 16:59:20 +0800 Subject: [PATCH 1/4] fix(lockservice): propagate cancellation through lock waits --- pkg/cnservice/server.go | 9 +- pkg/lockservice/deadlock.go | 19 +- pkg/lockservice/deadlock_test.go | 38 ++- pkg/lockservice/lock_table_keeper.go | 5 +- pkg/lockservice/lock_table_local.go | 18 +- pkg/lockservice/lock_table_local_test.go | 26 +- pkg/lockservice/lock_table_proxy.go | 49 ++-- pkg/lockservice/lock_table_proxy_test.go | 133 +++++++++++ pkg/lockservice/lock_table_remote.go | 56 +++-- pkg/lockservice/lock_table_remote_test.go | 185 +++++++++++++- pkg/lockservice/orphan_txn_test.go | 2 +- pkg/lockservice/service.go | 225 +++++++++++++----- pkg/lockservice/service_forward.go | 1 + pkg/lockservice/service_forward_test.go | 2 +- pkg/lockservice/service_observability.go | 34 ++- pkg/lockservice/service_observability_test.go | 3 +- pkg/lockservice/service_remote.go | 61 ++--- pkg/lockservice/service_remote_test.go | 39 +-- pkg/lockservice/service_test.go | 154 ++++++++++-- pkg/lockservice/test_helper.go | 3 +- pkg/lockservice/txn.go | 41 +++- pkg/lockservice/txn_test.go | 34 ++- pkg/lockservice/types.go | 2 +- pkg/lockservice/waiter.go | 30 +-- pkg/lockservice/waiter_test.go | 30 +++ 25 files changed, 961 insertions(+), 238 deletions(-) diff --git a/pkg/cnservice/server.go b/pkg/cnservice/server.go index 729044637b187..c0b86508e4887 100644 --- a/pkg/cnservice/server.go +++ b/pkg/cnservice/server.go @@ -808,8 +808,13 @@ func (s *service) initLockService() { cfg := s.getLockServiceConfig() s.lockService = lockservice.NewLockService( cfg, - lockservice.WithWait(func() { - <-s.hakeeperConnected + lockservice.WithWait(func(ctx context.Context) error { + select { + case <-s.hakeeperConnected: + return nil + case <-ctx.Done(): + return ctx.Err() + } })) runtime.ServiceRuntime(s.cfg.UUID).SetGlobalVariables(runtime.LockService, s.lockService) lockservice.SetLockServiceByServiceID(s.cfg.UUID, s.lockService) diff --git a/pkg/lockservice/deadlock.go b/pkg/lockservice/deadlock.go index 2e57fd2527263..a23305e870bcd 100644 --- a/pkg/lockservice/deadlock.go +++ b/pkg/lockservice/deadlock.go @@ -38,7 +38,7 @@ var ( type detector struct { logger *log.MOLogger c chan deadlockTxn - waitTxnsFetchFunc func(pb.WaitTxn, *waiters) (bool, error) + waitTxnsFetchFunc func(context.Context, pb.WaitTxn, *waiters) (bool, error) waitTxnAbortFunc func(pb.WaitTxn, error) ignoreTxns sync.Map // txnID -> any stopper *stopper.Stopper @@ -56,7 +56,7 @@ type detector struct { // txn. func newDeadlockDetector( logger *log.MOLogger, - waitTxnsFetchFunc func(pb.WaitTxn, *waiters) (bool, error), + waitTxnsFetchFunc func(context.Context, pb.WaitTxn, *waiters) (bool, error), waitTxnAbortFunc func(pb.WaitTxn, error), ) *detector { d := &detector{ @@ -82,6 +82,9 @@ func (d *detector) close() { d.mu.closed = true d.mu.Unlock() d.stopper.Stop() + d.mu.Lock() + clear(d.mu.activeCheckTxn) + d.mu.Unlock() close(d.c) } @@ -141,13 +144,16 @@ func (d *detector) doCheck(ctx context.Context) { w := &waiters{ignoreTxns: &d.ignoreTxns} for { + if ctx.Err() != nil { + return + } select { case <-ctx.Done(): return case txn := <-d.c: v2.TxnDeadlockDetectorQueueDepthGauge.Set(float64(len(d.c))) w.reset(txn) - hasDeadlock, deadlockTxn, err := d.checkDeadlock(w) + hasDeadlock, deadlockTxn, err := d.checkDeadlock(ctx, w) if hasDeadlock { if err == nil { err = ErrDeadLockDetected @@ -162,11 +168,14 @@ func (d *detector) doCheck(ctx context.Context) { } } -func (d *detector) checkDeadlock(w *waiters) (bool, pb.WaitTxn, error) { +func (d *detector) checkDeadlock(ctx context.Context, w *waiters) (bool, pb.WaitTxn, error) { for { + if err := ctx.Err(); err != nil { + return false, pb.WaitTxn{}, err + } // find deadlock txn := w.getCheckTargetTxn() - added, err := d.waitTxnsFetchFunc(txn, w) + added, err := d.waitTxnsFetchFunc(ctx, txn, w) if err != nil { logCheckDeadLockFailed(d.logger, txn, w.root.startTxn(), err) return false, pb.WaitTxn{}, err diff --git a/pkg/lockservice/deadlock_test.go b/pkg/lockservice/deadlock_test.go index 1b61a3134b131..85b92f3e40bd4 100644 --- a/pkg/lockservice/deadlock_test.go +++ b/pkg/lockservice/deadlock_test.go @@ -15,6 +15,7 @@ package lockservice import ( + "context" "encoding/hex" "testing" "time" @@ -23,6 +24,7 @@ import ( "github.com/matrixorigin/matrixone/pkg/common/runtime" pb "github.com/matrixorigin/matrixone/pkg/pb/lock" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestCheckWithDeadlock(t *testing.T) { @@ -43,7 +45,7 @@ func TestCheckWithDeadlock(t *testing.T) { d := newDeadlockDetector( runtime.DefaultRuntime().Logger(), - func(txn pb.WaitTxn, w *waiters) (bool, error) { + func(_ context.Context, txn pb.WaitTxn, w *waiters) (bool, error) { for _, v := range m[string(txn.TxnID)] { if !w.add(v, "") { return false, nil @@ -79,6 +81,34 @@ func TestCheckWithDeadlock(t *testing.T) { }) } +func TestDeadlockDetectorCloseCancelsCheck(t *testing.T) { + started := make(chan struct{}, 1) + aborted := make(chan struct{}, 1) + d := newDeadlockDetector( + runtime.DefaultRuntime().Logger(), + func(ctx context.Context, _ pb.WaitTxn, _ *waiters) (bool, error) { + select { + case started <- struct{}{}: + default: + } + <-ctx.Done() + return false, ctx.Err() + }, + func(pb.WaitTxn, error) { aborted <- struct{}{} }, + ) + require.NoError(t, d.check([]byte("holder"), pb.WaitTxn{TxnID: []byte("waiter")})) + <-started + d.close() + select { + case <-aborted: + t.Fatal("deadlock abort callback ran after detector cancellation") + default: + } + d.mu.Lock() + defer d.mu.Unlock() + require.Empty(t, d.mu.activeCheckTxn) +} + func TestCheckWithDeadlockWith2Txn(t *testing.T) { reuse.RunReuseTests(func() { txn1 := []byte("t1") @@ -95,7 +125,7 @@ func TestCheckWithDeadlockWith2Txn(t *testing.T) { d := newDeadlockDetector( runtime.DefaultRuntime().Logger(), - func(txn pb.WaitTxn, w *waiters) (bool, error) { + func(_ context.Context, txn pb.WaitTxn, w *waiters) (bool, error) { for _, v := range depends[string(txn.TxnID)] { if !w.add(v, "") { return false, nil @@ -282,7 +312,7 @@ func TestCheckWithComplexDeadlock(t *testing.T) { // Create the deadlock detector d := newDeadlockDetector( runtime.DefaultRuntime().Logger(), - func(txn pb.WaitTxn, w *waiters) (bool, error) { + func(_ context.Context, txn pb.WaitTxn, w *waiters) (bool, error) { for _, v := range depends[string(txn.TxnID)] { if !w.add(v, "") { return false, nil @@ -363,7 +393,7 @@ func TestCheckDeadlock(t *testing.T) { // Create the deadlock detector d := newDeadlockDetector( runtime.DefaultRuntime().Logger(), - func(txn pb.WaitTxn, w *waiters) (bool, error) { + func(_ context.Context, txn pb.WaitTxn, w *waiters) (bool, error) { for _, v := range depends[string(txn.TxnID)] { if !w.add(v, "") { return false, nil diff --git a/pkg/lockservice/lock_table_keeper.go b/pkg/lockservice/lock_table_keeper.go index f916deb7fdca5..3dfd3f0449bb1 100644 --- a/pkg/lockservice/lock_table_keeper.go +++ b/pkg/lockservice/lock_table_keeper.go @@ -266,10 +266,7 @@ func (k *lockTableKeeper) invalidateRemoteBind( } func (k *lockTableKeeper) doKeepLockTableBind(ctx context.Context) { - if k.service.isStatus(pb.Status_ServiceLockWaiting) && - k.service.activeTxnHolder.empty() { - k.service.setStatus(pb.Status_ServiceUnLockSucc) - } + k.service.tryCompleteDrain() req := acquireRequest() defer releaseRequest(req) diff --git a/pkg/lockservice/lock_table_local.go b/pkg/lockservice/lock_table_local.go index 6051675bbce00..6afdc8ecc2a90 100644 --- a/pkg/lockservice/lock_table_local.go +++ b/pkg/lockservice/lock_table_local.go @@ -390,23 +390,34 @@ func (l *localLockTable) unlock( } func (l *localLockTable) getLock( + ctx context.Context, key []byte, txn pb.WaitTxn, - fn func(Lock)) { + fn func(Lock)) error { + if err := ctx.Err(); err != nil { + return err + } l.mu.RLock() defer l.mu.RUnlock() + if err := ctx.Err(); err != nil { + return err + } if l.mu.closed { - return + return nil } lock, ok := l.mu.store.Get(key) if ok { fn(lock) } + return nil } func (l *localLockTable) getLockHolder(ctx context.Context, key []byte) (pb.WaitTxn, bool, error) { l.mu.RLock() defer l.mu.RUnlock() + if err := ctx.Err(); err != nil { + return pb.WaitTxn{}, false, err + } if l.mu.closed { return pb.WaitTxn{}, false, nil } @@ -454,6 +465,9 @@ func (l *localLockTable) doAcquireLock(c *lockContext) error { if l.mu.closed { return moerr.NewInvalidStateNoCtx("local lock table closed") } + if err := c.ctx.Err(); err != nil { + return err + } switch c.opts.Granularity { case pb.Granularity_Row: diff --git a/pkg/lockservice/lock_table_local_test.go b/pkg/lockservice/lock_table_local_test.go index c9a3a88652f56..e8f56a5e9ad24 100644 --- a/pkg/lockservice/lock_table_local_test.go +++ b/pkg/lockservice/lock_table_local_test.go @@ -109,7 +109,7 @@ func TestCloseLocalLockTableWithBlockedWaiter(t *testing.T) { require.Equal(t, ErrLockTableNotFound, err) }() - v, err := l.getLockTable(0, tableID) + v, err := l.getLockTable(context.Background(), 0, tableID) require.NoError(t, err) lt := v.(*localLockTable) for { @@ -408,7 +408,7 @@ func TestMergeRangeWithNoConflict(t *testing.T) { table := uint64(10) for _, c := range cases { stopper := stopper.NewStopper("") - v, err := l.getLockTableWithCreate(0, table, nil, pb.Sharding_None) + v, err := l.getLockTableWithCreate(context.Background(), 0, table, nil, pb.Sharding_None) require.NoError(t, err) lt := v.(*localLockTable) @@ -542,7 +542,7 @@ func TestLocalLockTableMultipleRowLocksCannotMissIfFoundSelfTxn(t *testing.T) { require.NoError(t, l.Unlock(ctx, []byte{2}, timestamp.Timestamp{})) wg.Wait() - v, err := l.getLockTable(0, tableID) + v, err := l.getLockTable(context.Background(), 0, tableID) require.NoError(t, err) lt := v.(*localLockTable) lt.mu.Lock() @@ -616,7 +616,7 @@ func TestIssue9856(t *testing.T) { json.MustUnmarshal([]byte(r), v) _, err := l.Lock(ctx, tableID, [][]byte{[]byte(v.Start), []byte(v.End)}, []byte("txn1"), option) require.NoError(t, err) - vv, err := l.getLockTable(0, tableID) + vv, err := l.getLockTable(context.Background(), 0, tableID) require.NoError(t, err) lt := vv.(*localLockTable) lt.mu.Lock() @@ -783,7 +783,7 @@ func TestLockedTSIsLastCommittedTS(t *testing.T) { defer cancel() tableID := uint64(10) - v, err := l.getLockTableWithCreate(0, tableID, nil, pb.Sharding_None) + v, err := l.getLockTableWithCreate(context.Background(), 0, tableID, nil, pb.Sharding_None) require.NoError(t, err) lt := v.(*localLockTable) lt.mu.Lock() @@ -854,7 +854,7 @@ func TestLockedTSIsLastCommittedTSWithRange(t *testing.T) { defer cancel() tableID := uint64(10) - v, err := l.getLockTableWithCreate(0, tableID, nil, pb.Sharding_None) + v, err := l.getLockTableWithCreate(context.Background(), 0, tableID, nil, pb.Sharding_None) require.NoError(t, err) lt := v.(*localLockTable) lt.mu.Lock() @@ -935,7 +935,7 @@ func Test15608(t *testing.T) { _, err := s1.Lock(ctx, table, rows, txn1, option) require.NoError(t, err, err) - v, err := s1.getLockTable(0, table) + v, err := s1.getLockTable(context.Background(), 0, table) require.NoError(t, err) lt := v.(*localLockTable) lt.options.beforeCloseFirstWaiter = func(c *lockContext) { @@ -1095,7 +1095,7 @@ func TestCannotHungIfRangeConflictWithRowMultiTimes(t *testing.T) { add(txn1, key4, pb.Granularity_Row) close(startTxn3) - v, err := l.getLockTable(0, tableID) + v, err := l.getLockTable(context.Background(), 0, tableID) require.NoError(t, err) lt := v.(*localLockTable) txn3WaitTimes := 0 @@ -1423,7 +1423,7 @@ func TestRangeLockModeUpgradeUpdatesBothEnds(t *testing.T) { require.NoError(t, err) // Verify both ends are Shared - v, err := l.getLockTable(0, tableID) + v, err := l.getLockTable(context.Background(), 0, tableID) require.NoError(t, err) lt := v.(*localLockTable) @@ -1527,7 +1527,7 @@ func TestSetModePairedRangeLockDirect(t *testing.T) { require.NoError(t, err) // Get the lock table and directly test setModePairedRangeLock - v, err := l.getLockTable(0, tableID) + v, err := l.getLockTable(context.Background(), 0, tableID) require.NoError(t, err) lt := v.(*localLockTable) @@ -1585,7 +1585,7 @@ func TestSetModePairedRangeLockFromRangeEnd(t *testing.T) { require.NoError(t, err) // Get the lock table and directly test setModePairedRangeLock from range-end - v, err := l.getLockTable(0, tableID) + v, err := l.getLockTable(context.Background(), 0, tableID) require.NoError(t, err) lt := v.(*localLockTable) @@ -1637,7 +1637,7 @@ func TestSetModePairedRangeLockRowLockNoOp(t *testing.T) { require.NoError(t, err) // Get the lock table - v, err := l.getLockTable(0, tableID) + v, err := l.getLockTable(context.Background(), 0, tableID) require.NoError(t, err) lt := v.(*localLockTable) @@ -1698,7 +1698,7 @@ func TestRangeLockWithInterleavedRowLocks(t *testing.T) { require.NoError(t, err) // Verify btree structure: [0:row] [1:range-start] [10:range-end] - v, err := l.getLockTable(0, tableID) + v, err := l.getLockTable(context.Background(), 0, tableID) require.NoError(t, err) lt := v.(*localLockTable) diff --git a/pkg/lockservice/lock_table_proxy.go b/pkg/lockservice/lock_table_proxy.go index 416717dc2ee32..4eb394b2e09a8 100644 --- a/pkg/lockservice/lock_table_proxy.go +++ b/pkg/lockservice/lock_table_proxy.go @@ -72,6 +72,11 @@ func (lp *localLockTableProxy) lock( } lp.mu.Lock() + if err := ctx.Err(); err != nil { + lp.mu.Unlock() + cb(pb.Result{}, err) + return + } key := util.UnsafeBytesToString(rows[0]) if _, ok := lp.mu.pendingLastHolderUnlocks[key]; ok { // The owner may already have applied the last-holder Unlock even @@ -93,7 +98,6 @@ func (lp *localLockTableProxy) lock( } first := v.isEmpty() - r := v.result w := v.add( lp.serviceID, txn, @@ -122,18 +126,29 @@ func (lp *localLockTableProxy) lock( return } - defer func() { - bind := lp.getBind() - err := txn.lockAdded(bind.Group, bind, rows, lp.logger) - cb(r, err) - }() - // wait first done if w != nil { - w.wait(ctx, lp.logger) - return + value := w.wait(ctx, lp.logger) + if value.err != nil { + lp.mu.Lock() + v.remove(txn) + lp.mu.Unlock() + cb(pb.Result{}, value.err) + return + } } + lp.mu.Lock() + r := v.result + lp.mu.Unlock() + bind := lp.getBind() + err := txn.lockAdded(bind.Group, bind, rows, lp.logger) + if err != nil { + lp.mu.Lock() + v.remove(txn) + lp.mu.Unlock() + } + cb(r, err) } func (lp *localLockTableProxy) unlock( @@ -314,10 +329,11 @@ func (lp *localLockTableProxy) isPendingRemoteHolderLocked( } func (lp *localLockTableProxy) getLock( + ctx context.Context, key []byte, txn pb.WaitTxn, - fn func(Lock)) { - lp.remote.getLock(key, txn, fn) + fn func(Lock)) error { + return lp.remote.getLock(ctx, key, txn, fn) } func (lp *localLockTableProxy) getLockHolder(ctx context.Context, key []byte) (pb.WaitTxn, bool, error) { @@ -347,9 +363,10 @@ func (s *sharedOps) done( logger *log.MOLogger, ) { for idx, cb := range s.cbs { - cb(r, err) - if idx > 0 { - s.waiters[idx].notify(notifyValue{}, logger) + if idx == 0 && cb != nil { + cb(r, err) + } else if s.waiters[idx] != nil { + s.waiters[idx].notify(notifyValue{err: err}, logger) } s.cbs[idx] = nil s.waiters[idx] = nil @@ -394,6 +411,10 @@ func (s *sharedOps) add( v := txn.toWaitTxn(serviceID, true) w = acquireWaiter(v, "share ops add", logger) w.setStatus(blocking) + // The waiting goroutine owns its callback. sharedOps.done only publishes + // the first remote result through the waiter, so completion and caller + // cancellation have one waiter-status linearization point. + cb = nil } if hasHolder { cb = nil diff --git a/pkg/lockservice/lock_table_proxy_test.go b/pkg/lockservice/lock_table_proxy_test.go index 749abc6de3887..1988537dd4b6c 100644 --- a/pkg/lockservice/lock_table_proxy_test.go +++ b/pkg/lockservice/lock_table_proxy_test.go @@ -17,6 +17,7 @@ package lockservice import ( "context" "errors" + "sync/atomic" "testing" "time" @@ -44,6 +45,37 @@ type recordingUnlockTable struct { mutations []pb.ExtraMutation } +type blockingProxyLockTable struct { + lockTable + bind pb.LockTable + started chan struct{} + release chan struct{} + err error +} + +func (t *blockingProxyLockTable) lock( + ctx context.Context, + _ *activeTxn, + _ [][]byte, + _ LockOptions, + cb func(pb.Result, error), +) { + select { + case t.started <- struct{}{}: + default: + } + select { + case <-t.release: + cb(pb.Result{}, t.err) + case <-ctx.Done(): + cb(pb.Result{}, ctx.Err()) + } +} + +func (t *blockingProxyLockTable) getBind() pb.LockTable { + return t.bind +} + func (t *recordingUnlockTable) unlockWithContext( ctx context.Context, txn *activeTxn, @@ -86,6 +118,107 @@ func (t *unlockBeforeApplyErrorTable) unlockWithContext( return t.err } +func TestProxySharedLockCancellationWhileFirstRemoteLockInFlight(t *testing.T) { + for _, remoteErr := range []error{nil, errors.New("remote lock failed")} { + name := "remote-success" + if remoteErr != nil { + name = "remote-failure" + } + t.Run(name, func(t *testing.T) { + bind := pb.LockTable{ + Group: 0, + Table: 1, + ServiceID: "remote", + Valid: true, + } + remote := &blockingProxyLockTable{ + bind: bind, + started: make(chan struct{}, 1), + release: make(chan struct{}), + err: remoteErr, + } + proxy := newLockTableProxy("local", remote, getLogger("")).(*localLockTableProxy) + rows := [][]byte{[]byte("row")} + options := LockOptions{LockOptions: newTestRowSharedOptions()} + firstTxn := newActiveTxn([]byte("first"), "first", newFixedSlicePool(4), "") + secondTxn := newActiveTxn([]byte("second"), "second", newFixedSlicePool(4), "") + + firstDone := make(chan error, 1) + go func() { + firstTxn.Lock() + defer firstTxn.Unlock() + proxy.lock(context.Background(), firstTxn, rows, options, func(_ pb.Result, err error) { + firstDone <- err + }) + }() + select { + case <-remote.started: + case <-time.After(time.Second): + t.Fatal("first remote lock did not start") + } + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + var secondCallbacks atomic.Int32 + var secondLockAdded atomic.Int32 + secondTxn.beforeLockAdded = func([]byte, [][]byte) error { + secondLockAdded.Add(1) + return nil + } + secondDone := make(chan error, 2) + go func() { + secondTxn.Lock() + defer secondTxn.Unlock() + proxy.lock(ctx, secondTxn, rows, options, func(_ pb.Result, err error) { + secondCallbacks.Add(1) + secondDone <- err + }) + }() + + require.Eventually(t, func() bool { + proxy.mu.Lock() + defer proxy.mu.Unlock() + ops := proxy.mu.holders[string(rows[0])] + return ops != nil && len(ops.txns) == 2 && ops.waiters[1] != nil + }, time.Second, time.Millisecond, "second shared lock was not admitted") + + cancel() + select { + case err := <-secondDone: + require.ErrorIs(t, err, context.Canceled) + case <-time.After(time.Second): + t.Fatal("second shared lock ignored cancellation") + } + require.Equal(t, int32(1), secondCallbacks.Load()) + require.Zero(t, secondLockAdded.Load()) + proxy.mu.Lock() + ops := proxy.mu.holders[string(rows[0])] + require.Len(t, ops.txns, 1) + require.Same(t, firstTxn, ops.txns[0]) + proxy.mu.Unlock() + + close(remote.release) + select { + case err := <-firstDone: + if remoteErr == nil { + require.NoError(t, err) + } else { + require.ErrorIs(t, err, remoteErr) + } + case <-time.After(time.Second): + t.Fatal("first remote lock did not finish") + } + require.Equal(t, int32(1), secondCallbacks.Load()) + require.Zero(t, secondLockAdded.Load()) + select { + case err := <-secondDone: + t.Fatalf("second callback ran more than once: %v", err) + default: + } + }) + } +} + func TestProxySharedLock(t *testing.T) { runLockServiceTests( t, diff --git a/pkg/lockservice/lock_table_remote.go b/pkg/lockservice/lock_table_remote.go index 44d72663ce5c3..63ec90bae239b 100644 --- a/pkg/lockservice/lock_table_remote.go +++ b/pkg/lockservice/lock_table_remote.go @@ -187,7 +187,7 @@ func (l *remoteLockTable) lock( // swallows the error, the transaction will not be abort. originalErr := err txn.Unlock() - e := l.handleError(err, true) + e := l.handleErrorWithContext(ctx, err, true) txn.Lock() if !bytes.Equal(req.Lock.TxnID, txn.txnID) { cb(pb.Result{}, ErrTxnNotFound) @@ -273,25 +273,42 @@ func (l *remoteLockTable) unlockWithContext( } func (l *remoteLockTable) getLock( + ctx context.Context, key []byte, txn pb.WaitTxn, - fn func(Lock)) { + fn func(Lock)) error { + if err := ctx.Err(); err != nil { + return err + } backoff := remoteRetryInitialBackoff for { - lock, ok, err := l.doGetLock(key, txn) + if err := ctx.Err(); err != nil { + return err + } + lock, ok, err := l.doGetLock(ctx, key, txn) if err == nil { + if ok { + defer lock.close(notifyValue{}) + } + if err := ctx.Err(); err != nil { + return err + } if ok { fn(lock) - lock.close(notifyValue{}) } - return + return nil } // why use loop is similar to unlock - if err = l.handleError(err, false); err == nil { - return + if err = l.handleErrorWithContext(ctx, err, false); err == nil { + // The bind-change handler replaces this table in service.tableGroups. + // Let the caller reacquire it instead of treating the stale snapshot as + // an empty waiting list. + return ErrLockTableBindChanged + } + if err := waitRemoteRetryBackoffWithContext(ctx, backoff); err != nil { + return err } - waitRemoteRetryBackoff(backoff) backoff = nextRemoteRetryBackoff(backoff) } } @@ -309,7 +326,7 @@ func (l *remoteLockTable) getLockHolder(ctx context.Context, key []byte) (pb.Wai if err := ctx.Err(); err != nil { return pb.WaitTxn{}, false, err } - if err = l.handleError(err, false); err == nil { + if err = l.handleErrorWithContext(ctx, err, false); err == nil { // The bind-change handler replaces the lock-table object in service.tableGroups. // This in-flight remote table still carries the stale bind, so let the service // reacquire the current table before retrying the holder lookup. @@ -378,8 +395,8 @@ func (l *remoteLockTable) doUnlock( return moerr.AttachCause(ctx, err) } -func (l *remoteLockTable) doGetLock(key []byte, txn pb.WaitTxn) (Lock, bool, error) { - ctx, cancel := context.WithTimeoutCause(context.Background(), defaultRPCTimeout, moerr.CauseDoGetLock) +func (l *remoteLockTable) doGetLock(parent context.Context, key []byte, txn pb.WaitTxn) (Lock, bool, error) { + ctx, cancel := context.WithTimeoutCause(parent, defaultRPCTimeout, moerr.CauseDoGetLock) defer cancel() req := acquireRequest() @@ -454,17 +471,6 @@ func (l *remoteLockTable) close(reason closeReason) { logLockTableClosed(l.logger, l.bind, true, reason) } -func (l *remoteLockTable) handleError( - err error, - mustHandleLockBindChangedErr bool, -) error { - return l.handleErrorWithContext( - context.Background(), - err, - mustHandleLockBindChangedErr, - ) -} - func (l *remoteLockTable) handleErrorWithContext( ctx context.Context, err error, @@ -501,6 +507,9 @@ func (l *remoteLockTable) handleErrorWithContext( ) if err != nil { logGetRemoteBindFailed(l.logger, l.bind.Table, err) + if ctxErr := ctx.Err(); ctxErr != nil { + return ctxErr + } return oldError } if new.Changed(l.bind) { @@ -559,6 +568,9 @@ func (l *remoteLockTable) maybeHandleBindChanged( ) if err != nil { logGetRemoteBindFailed(l.logger, l.bind.Table, err) + if ctxErr := ctx.Err(); ctxErr != nil { + return ctxErr + } return ErrLockTableBindChanged } if !refreshedBind.Changed(l.bind) { diff --git a/pkg/lockservice/lock_table_remote_test.go b/pkg/lockservice/lock_table_remote_test.go index ba4938e5e8d44..eb4277d30e028 100644 --- a/pkg/lockservice/lock_table_remote_test.go +++ b/pkg/lockservice/lock_table_remote_test.go @@ -61,7 +61,7 @@ type blockingBindRefreshClient struct { func (c *blockingBindRefreshClient) Send(ctx context.Context, req *pb.Request) (*pb.Response, error) { switch req.Method { - case pb.Method_Unlock: + case pb.Method_Unlock, pb.Method_GetTxnLock: return nil, io.ErrUnexpectedEOF case pb.Method_GetBind: select { @@ -81,6 +81,176 @@ func (c *blockingBindRefreshClient) AsyncSend(context.Context, *pb.Request) (*mo func (c *blockingBindRefreshClient) Close() error { return nil } +type blockingGetLockClient struct { + started chan struct{} +} + +func (c *blockingGetLockClient) Send(ctx context.Context, req *pb.Request) (*pb.Response, error) { + if req.Method != pb.Method_GetTxnLock { + return nil, io.ErrClosedPipe + } + select { + case c.started <- struct{}{}: + default: + } + <-ctx.Done() + return nil, ctx.Err() +} + +func (c *blockingGetLockClient) AsyncSend(context.Context, *pb.Request) (*morpc.Future, error) { + return nil, io.ErrClosedPipe +} + +func (c *blockingGetLockClient) Close() error { return nil } + +type retryingGetLockClient struct { + bind pb.LockTable + started chan struct{} + release chan struct{} +} + +func (c *retryingGetLockClient) Send(_ context.Context, req *pb.Request) (*pb.Response, error) { + switch req.Method { + case pb.Method_GetTxnLock: + select { + case c.started <- struct{}{}: + default: + } + select { + case <-c.release: + return &pb.Response{}, nil + default: + return nil, io.ErrUnexpectedEOF + } + case pb.Method_GetBind: + resp := &pb.Response{} + resp.GetBind.LockTable = c.bind + resp.GetBind.AllocatorID = c.bind.AllocatorID + resp.GetBind.AllocatorVersion = c.bind.Version + return resp, nil + default: + return nil, io.ErrClosedPipe + } +} + +func (c *retryingGetLockClient) AsyncSend(context.Context, *pb.Request) (*morpc.Future, error) { + return nil, io.ErrClosedPipe +} + +func (c *retryingGetLockClient) Close() error { return nil } + +func TestRemoteGetLockWithContextStopsOnCancellation(t *testing.T) { + client := &blockingGetLockClient{started: make(chan struct{}, 1)} + remote := newRemoteLockTable( + "s1", + time.Second, + pb.LockTable{ServiceID: "s2", Table: 1, Valid: true}, + client, + func(pb.LockTable) {}, + getLogger(""), + ) + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + called := false + go func() { + done <- remote.getLock(ctx, []byte("row"), pb.WaitTxn{TxnID: []byte("txn")}, func(Lock) { + called = true + }) + }() + <-client.started + cancel() + require.ErrorIs(t, <-done, context.Canceled) + require.False(t, called) +} + +func TestRemoteGetLockPreservesCancellationDuringBindRefresh(t *testing.T) { + client := &blockingBindRefreshClient{bindRefreshStarted: make(chan struct{}, 1)} + remote := newRemoteLockTable( + "s1", + time.Second, + pb.LockTable{ServiceID: "s2", Table: 1, OriginTable: 1, Valid: true}, + client, + func(pb.LockTable) {}, + getLogger(""), + ) + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { + done <- remote.getLock(ctx, []byte("row"), pb.WaitTxn{TxnID: []byte("txn")}, func(Lock) {}) + }() + <-client.bindRefreshStarted + cancel() + require.ErrorIs(t, <-done, context.Canceled) +} + +func TestDeadlockDetectorCloseCancelsRemoteGetLockRetry(t *testing.T) { + bind := pb.LockTable{ + Group: 0, + Table: 1, + OriginTable: 1, + ServiceID: "s2", + Version: 1, + Valid: true, + AllocatorID: "allocator-1", + } + client := &retryingGetLockClient{ + bind: bind, + started: make(chan struct{}, 1), + release: make(chan struct{}), + } + remote := newRemoteLockTable( + "s1", + time.Second, + bind, + client, + func(pb.LockTable) {}, + getLogger(""), + ) + callbackCalled := make(chan struct{}, 1) + abortCalled := make(chan struct{}, 1) + d := newDeadlockDetector( + getLogger(""), + func(ctx context.Context, _ pb.WaitTxn, _ *waiters) (bool, error) { + err := remote.getLock(ctx, []byte("row"), pb.WaitTxn{TxnID: []byte("holder")}, func(Lock) { + callbackCalled <- struct{}{} + }) + return err == nil, err + }, + func(pb.WaitTxn, error) { abortCalled <- struct{}{} }, + ) + require.NoError(t, d.check([]byte("holder"), pb.WaitTxn{TxnID: []byte("waiter")})) + <-client.started + + closed := make(chan struct{}) + go func() { + d.close() + close(closed) + }() + + returnedPromptly := false + select { + case <-closed: + returnedPromptly = true + case <-time.After(time.Second): + close(client.release) + <-closed + } + require.True(t, returnedPromptly, "detector close remained blocked in remote GetTxnLock retry") + select { + case <-callbackCalled: + t.Fatal("remote lock callback ran after detector cancellation") + default: + } + select { + case <-abortCalled: + t.Fatal("deadlock abort callback ran after detector cancellation") + default: + } + d.mu.Lock() + defer d.mu.Unlock() + require.Empty(t, d.mu.activeCheckTxn) +} + func TestRemoteUnlockWithContextStopsOnCancellation(t *testing.T) { defer leaktest.AfterTest(t)() @@ -198,7 +368,7 @@ func TestRemoteNewBindRefreshHonorsContext(t *testing.T) { cancel() select { case err := <-done: - require.ErrorIs(t, err, ErrLockTableBindChanged) + require.ErrorIs(t, err, context.Canceled) case <-time.After(time.Second): require.FailNow(t, "remote bind refresh ignored cancellation") } @@ -370,8 +540,7 @@ func TestLockRemoteWithContextTimeoutTracksLockForUnlock(t *testing.T) { }() l.lock(ctx, txn, [][]byte{{1}}, LockOptions{}, func(r pb.Result, err error) { - require.Error(t, err) - require.True(t, moerr.IsMoErrCode(err, moerr.ErrBackendCannotConnect)) + require.ErrorIs(t, err, context.DeadlineExceeded) }) holder := txn.getHoldLocksLocked(l.bind.Group) require.Contains(t, holder.tableKeys, l.bind.Table) @@ -898,7 +1067,7 @@ func TestRemoteBindRefreshRejectsSupersededAllocatorBind(t *testing.T) { l.allocatorStateProvider = svc.allocatorStateSnapshot l.allocatorBindChangedHandler = svc.handleBindChangedFromAllocator - err := l.handleError(moerr.NewRPCTimeoutNoCtx(), true) + err := l.handleErrorWithContext(context.Background(), moerr.NewRPCTimeoutNoCtx(), true) require.True(t, moerr.IsMoErrCode(err, moerr.ErrLockTableBindChanged)) require.Nil(t, svc.tableGroups.get(oldBind.Group, oldBind.Table)) require.Equal(t, newAllocator.id, svc.lastAllocatorID) @@ -972,7 +1141,7 @@ func TestGetLockRemoteWithRetry(t *testing.T) { ) }, func(l *remoteLockTable, s Server) { - l.getLock([]byte("row1"), pb.WaitTxn{TxnID: []byte("txn1")}, func(lock Lock) { + _ = l.getLock(context.Background(), []byte("row1"), pb.WaitTxn{TxnID: []byte("txn1")}, func(lock Lock) { called = true assert.Equal(t, byte(pb.Granularity_Row), lock.value) }) @@ -1177,7 +1346,9 @@ func TestRemoteWithBindChanged(t *testing.T) { l.unlock(txn, nil, timestamp.Timestamp{}) assert.Equal(t, newBind, <-c) - l.getLock(txnID, pb.WaitTxn{TxnID: []byte{1}}, nil) + require.ErrorIs(t, + l.getLock(context.Background(), txnID, pb.WaitTxn{TxnID: []byte{1}}, nil), + ErrLockTableBindChanged) assert.Equal(t, newBind, <-c) reuse.Free(txn, nil) }, diff --git a/pkg/lockservice/orphan_txn_test.go b/pkg/lockservice/orphan_txn_test.go index 8795ef1f07580..27675d803847d 100644 --- a/pkg/lockservice/orphan_txn_test.go +++ b/pkg/lockservice/orphan_txn_test.go @@ -865,7 +865,7 @@ func TestOrphanTxnHolderCanBeRelease(t *testing.T) { require.NoError(t, s2.Unlock(ctx, txn2, timestamp.Timestamp{})) close(ch) - v, err := s1.getLockTable(0, table) + v, err := s1.getLockTable(context.Background(), 0, table) require.NoError(t, err) lt := v.(*localLockTable) diff --git a/pkg/lockservice/service.go b/pkg/lockservice/service.go index 787820db7a944..d074f21fe4fdb 100644 --- a/pkg/lockservice/service.go +++ b/pkg/lockservice/service.go @@ -40,7 +40,7 @@ import ( ) // WithWait setup wait func to wait some condition ready -func WithWait(wait func()) Option { +func WithWait(wait func(context.Context) error) Option { return func(s *service) { s.option.wait = wait } @@ -77,15 +77,16 @@ type service struct { mu struct { sync.RWMutex - restartTime timestamp.Timestamp - status pb.Status - groupTables [][]pb.LockTable - lockTableRef map[uint32]map[uint64]uint64 - allocating map[uint32]map[uint64]chan struct{} + restartTime timestamp.Timestamp + status pb.Status + lockAdmissions uint64 + groupTables [][]pb.LockTable + lockTableRef map[uint32]map[uint64]uint64 + allocating map[uint32]map[uint64]chan struct{} } option struct { - wait func() + wait func(context.Context) error beforeRemoteLockBindCheck func() serverOpts []ServerOption } @@ -157,10 +158,14 @@ func (s *service) Lock( rows [][]byte, txnID []byte, options pb.LockOptions) (pb.Result, error) { + if err := ctx.Err(); err != nil { + return pb.Result{}, err + } - if !s.canLockOnServiceStatus(txnID, options, tableID, rows) { + if !s.beginLockAdmission(txnID, options, tableID, rows) { return pb.Result{}, moerr.NewNewTxnInCNRollingRestart() } + defer s.endLockAdmission() v2.TxnLockTotalCounter.Inc() options.Validate(rows) @@ -170,7 +175,12 @@ func (s *service) Lock( v2.TxnAcquireLockDurationHistogram.Observe(time.Since(start).Seconds()) }() - s.wait() + if err := s.wait(ctx); err != nil { + return pb.Result{}, err + } + if err := ctx.Err(); err != nil { + return pb.Result{}, err + } // FIXME(fagongzi): too many mem alloc in trace ctx, span := trace.Debug(ctx, "lockservice.lock") @@ -180,11 +190,14 @@ func (s *service) Lock( return s.forwardLock(ctx, tableID, rows, txnID, options) } - txn := s.activeTxnHolder.getActiveTxn(txnID, true, "") - l, err := s.getLockTableWithCreate(options.Group, tableID, rows, options.Sharding) + l, err := s.getLockTableWithCreate(ctx, options.Group, tableID, rows, options.Sharding) if err != nil { return pb.Result{}, err } + if err := ctx.Err(); err != nil { + return pb.Result{}, err + } + txn := s.activeTxnHolder.getActiveTxn(txnID, true, "") s.bindChangeMu.RLock() // All txn lock op must be serial. And avoid dead lock between doAcquireLock @@ -206,6 +219,11 @@ func (s *service) Lock( s.bindChangeMu.RUnlock() return pb.Result{}, ErrLockTableBindChanged } + if err := ctx.Err(); err != nil { + txn.Unlock() + s.bindChangeMu.RUnlock() + return pb.Result{}, err + } // it needs to inc table bind ref when set restart cn bind := l.getBind() @@ -274,7 +292,9 @@ func (s *service) unlockUnknownCommit( if err := ctx.Err(); err != nil { return err } - s.wait() + if err := s.wait(ctx); err != nil { + return err + } if err := ctx.Err(); err != nil { return err } @@ -300,7 +320,9 @@ func (s *service) unlockUnknownCommit( ctx, txnID, commitTS, - s.getLockTable, + func(group uint32, table uint64) (lockTable, error) { + return s.getLockTable(ctx, group, table) + }, s.logger, mutations..., ); err != nil { @@ -312,9 +334,7 @@ func (s *service) unlockUnknownCommit( } if !s.isStatus(pb.Status_ServiceLockEnable) { s.reduceCanMoveGroupTables(txn) - if s.isStatus(pb.Status_ServiceLockWaiting) && s.activeTxnHolder.empty() { - s.setStatus(pb.Status_ServiceUnLockSucc) - } + s.tryCompleteDrain() } s.deadlockDetector.txnClosed(txnID) reuse.Free(txn, nil) @@ -334,7 +354,9 @@ func (s *service) unlockWithContext( if err := ctx.Err(); err != nil { return err } - s.wait() + if err := s.wait(ctx); err != nil { + return err + } if err := ctx.Err(); err != nil { return err } @@ -352,14 +374,13 @@ func (s *service) unlockWithContext( if !s.isStatus(pb.Status_ServiceLockEnable) { s.reduceCanMoveGroupTables(txn) - if s.isStatus(pb.Status_ServiceLockWaiting) && - s.activeTxnHolder.empty() { - s.setStatus(pb.Status_ServiceUnLockSucc) - } + s.tryCompleteDrain() } defer logUnlockTxn(s.logger, txn)() - err := txn.closeWithContext(ctx, txnID, commitTS, s.getLockTable, s.logger, mutations...) + err := txn.closeWithContext(ctx, txnID, commitTS, func(group uint32, table uint64) (lockTable, error) { + return s.getLockTable(ctx, group, table) + }, s.logger, mutations...) // The deadlock detector will hold the deadlocked transaction that is aborted // to avoid the situation where the deadlock detection is interfered with by // the abort transaction. When a transaction is unlocked, the deadlock detector @@ -437,7 +458,7 @@ func (s *service) reduceCanMoveGroupTables(txn *activeTxn) { func (s *service) checkCanMoveGroupTables() { s.mu.Lock() defer s.mu.Unlock() - if s.mu.status != pb.Status_ServiceLockEnable { + if s.mu.status != pb.Status_ServiceLockEnable || s.mu.lockAdmissions > 0 { return } @@ -469,12 +490,13 @@ func (s *service) incRef(group uint32, table uint64) { s.mu.lockTableRef[group][table]++ } -func (s *service) canLockOnServiceStatus( +func (s *service) canLockOnServiceStatusLocked( txnID []byte, opts pb.LockOptions, tableID uint64, - rows [][]byte) bool { - if s.isStatus(pb.Status_ServiceLockEnable) { + rows [][]byte, +) bool { + if s.mu.status == pb.Status_ServiceLockEnable { return true } if opts.Sharding == pb.Sharding_ByRow { @@ -483,19 +505,60 @@ func (s *service) canLockOnServiceStatus( if s.activeTxnHolder.hasActiveTxn(txnID) { return true } - if !s.validGroupTable(opts.Group, tableID) { + if _, ok := s.mu.lockTableRef[opts.Group][tableID]; !ok { logCanLockOnService(s.logger, s.serviceID) return false } if s.activeTxnHolder.empty() { return false } - if opts.SnapShotTs.LessEq(s.getRestartTime()) { + if opts.SnapShotTs.LessEq(s.mu.restartTime) { return true } return false } +func (s *service) beginLockAdmission( + txnID []byte, + opts pb.LockOptions, + tableID uint64, + rows [][]byte, +) bool { + s.mu.Lock() + defer s.mu.Unlock() + if !s.canLockOnServiceStatusLocked(txnID, opts, tableID, rows) { + return false + } + s.mu.lockAdmissions++ + return true +} + +func (s *service) endLockAdmission() { + s.mu.Lock() + defer s.mu.Unlock() + if s.mu.lockAdmissions == 0 { + panic("lock admission underflow") + } + s.mu.lockAdmissions-- + s.tryCompleteDrainLocked() +} + +func (s *service) tryCompleteDrain() { + s.mu.Lock() + defer s.mu.Unlock() + s.tryCompleteDrainLocked() +} + +func (s *service) tryCompleteDrainLocked() { + if s.mu.status != pb.Status_ServiceLockWaiting || + s.mu.lockAdmissions != 0 || + !s.activeTxnHolder.empty() { + return + } + logStatusChange(s.logger, s.mu.status, pb.Status_ServiceUnLockSucc) + s.mu.status = pb.Status_ServiceUnLockSucc +} + func (s *service) validGroupTable(group uint32, tableID uint64) bool { s.mu.RLock() defer s.mu.RUnlock() @@ -503,12 +566,6 @@ func (s *service) validGroupTable(group uint32, tableID uint64) bool { return ok } -func (s *service) getRestartTime() timestamp.Timestamp { - s.mu.RLock() - defer s.mu.RUnlock() - return s.mu.restartTime -} - func (s *service) GetServiceID() string { return s.serviceID } @@ -582,7 +639,10 @@ func (s *service) isStatus(status pb.Status) bool { return s.mu.status == status } -func (s *service) fetchTxnWaitingList(txn pb.WaitTxn, waiters *waiters) (bool, error) { +func (s *service) fetchTxnWaitingList(ctx context.Context, txn pb.WaitTxn, waiters *waiters) (bool, error) { + if err := ctx.Err(); err != nil { + return false, err + } if txn.CreatedOn == s.serviceID { activeTxn := s.activeTxnHolder.getActiveTxn(txn.TxnID, false, "") // the active txn closed @@ -594,13 +654,16 @@ func (s *service) fetchTxnWaitingList(txn pb.WaitTxn, waiters *waiters) (bool, e return true, nil } return activeTxn.fetchWhoWaitingMe( + ctx, s.serviceID, txnID, waiters.add, - s.getLockTable), nil + func(ctx context.Context, group uint32, table uint64) (lockTable, error) { + return s.getLockTable(ctx, group, table) + }) } - waitingList, err := s.getTxnWaitingListOnRemote(txn.TxnID, txn.CreatedOn) + waitingList, err := s.getTxnWaitingListOnRemote(ctx, txn.TxnID, txn.CreatedOn) if err != nil { return false, err } @@ -629,16 +692,14 @@ func (s *service) abortDeadlockTxn(wait pb.WaitTxn, err error) { activeTxn.abort(wait, err, s.logger) } -func (s *service) getLockTable( - group uint32, - tableID uint64) (lockTable, error) { +func (s *service) getLockTable(ctx context.Context, group uint32, tableID uint64) (lockTable, error) { + if err := ctx.Err(); err != nil { + return nil, err + } if v := s.tableGroups.get(group, tableID); v != nil { return v, nil } - return s.waitLockTableBind( - group, - tableID, - false), nil + return s.waitLockTableBind(ctx, group, tableID, false) } func (s *service) getAllocatingC( @@ -656,27 +717,37 @@ func (s *service) getAllocatingC( } func (s *service) waitLockTableBind( + ctx context.Context, group uint32, tableID uint64, - locked bool) lockTable { + locked bool) (lockTable, error) { c := s.getAllocatingC(group, tableID, locked) if c != nil { - <-c + select { + case <-c: + case <-ctx.Done(): + return nil, ctx.Err() + } } - return s.tableGroups.get(group, tableID) + if err := ctx.Err(); err != nil { + return nil, err + } + return s.tableGroups.get(group, tableID), nil } -func (s *service) getLockTableWithCreate( - group uint32, - tableID uint64, - rows [][]byte, - sharding pb.Sharding) (lockTable, error) { +func (s *service) getLockTableWithCreate(ctx context.Context, group uint32, tableID uint64, rows [][]byte, sharding pb.Sharding) (lockTable, error) { + if err := ctx.Err(); err != nil { + return nil, err + } originTableID := tableID if sharding == pb.Sharding_ByRow { tableID = ShardingByRow(rows[0]) } if v := s.tableGroups.get(group, tableID); v != nil { + if err := ctx.Err(); err != nil { + return nil, err + } return v, nil } @@ -686,9 +757,17 @@ func (s *service) getLockTableWithCreate( waitC := s.getAllocatingC(group, tableID, true) if waitC != nil { s.mu.Unlock() - <-waitC + select { + case <-waitC: + case <-ctx.Done(): + return nil + } s.mu.Lock() } + if err := ctx.Err(); err != nil { + s.mu.Unlock() + return nil + } v := s.tableGroups.get(group, tableID) if v == nil { @@ -704,19 +783,25 @@ func (s *service) getLockTableWithCreate( return v } - if v := fn(); v != nil { + v := fn() + if c != nil { + defer func() { + s.mu.Lock() + defer s.mu.Unlock() + delete(s.mu.allocating[group], tableID) + close(c) + }() + } + if v != nil { return v, nil } - - defer func() { - s.mu.Lock() - defer s.mu.Unlock() - delete(s.mu.allocating[group], tableID) - close(c) - }() + if err := ctx.Err(); err != nil { + return nil, err + } requestAllocator := s.allocatorStateSnapshot() - bind, allocator, err := getLockTableBind( + bind, allocator, err := getLockTableBindWithContext( + ctx, s.remote.client, group, tableID, @@ -726,8 +811,12 @@ func (s *service) getLockTableWithCreate( if err != nil { return nil, err } + if err := ctx.Err(); err != nil { + return nil, err + } return s.publishLockTableBindFromAllocator( + ctx, "get-bind", group, tableID, @@ -737,6 +826,7 @@ func (s *service) getLockTableWithCreate( } func (s *service) publishLockTableBindFromAllocator( + ctx context.Context, source string, group uint32, tableID uint64, @@ -746,7 +836,12 @@ func (s *service) publishLockTableBindFromAllocator( ) (lockTable, error) { s.allocatorVersionMu.Lock() defer s.allocatorVersionMu.Unlock() + if err := ctx.Err(); err != nil { + return nil, err + } + // Allocator-state observation and bind publication form one non-cancellable + // state transition. Once it starts, finish it and return its actual result. if _, accepted := s.observeAllocatorStateLocked( source, allocator, @@ -1120,11 +1215,11 @@ func (s *service) createLockTableByBind(bind pb.LockTable) lockTable { } } -func (s *service) wait() { +func (s *service) wait(ctx context.Context) error { if s.option.wait == nil { - return + return nil } - s.option.wait() + return s.option.wait(ctx) } type activeTxnHolder interface { diff --git a/pkg/lockservice/service_forward.go b/pkg/lockservice/service_forward.go index 8c7503b06afce..c77bb6a156c33 100644 --- a/pkg/lockservice/service_forward.go +++ b/pkg/lockservice/service_forward.go @@ -27,6 +27,7 @@ func (s *service) forwardLock( txnID []byte, opts pb.LockOptions) (pb.Result, error) { l, err := s.getLockTableWithCreate( + ctx, opts.Group, tableID, rows, diff --git a/pkg/lockservice/service_forward_test.go b/pkg/lockservice/service_forward_test.go index f009b744be3dc..4791fd21b7da0 100644 --- a/pkg/lockservice/service_forward_test.go +++ b/pkg/lockservice/service_forward_test.go @@ -37,7 +37,7 @@ func TestForwardLock(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), time.Second*10) defer cancel() - _, err := l2.getLockTableWithCreate(0, tableID, nil, pb.Sharding_None) + _, err := l2.getLockTableWithCreate(context.Background(), 0, tableID, nil, pb.Sharding_None) require.NoError(t, err) txn1 := []byte("txn1") diff --git a/pkg/lockservice/service_observability.go b/pkg/lockservice/service_observability.go index dc0906a248b27..65cdf1d46c2cb 100644 --- a/pkg/lockservice/service_observability.go +++ b/pkg/lockservice/service_observability.go @@ -24,6 +24,9 @@ import ( func (s *service) GetWaitingList( ctx context.Context, txnID []byte) (bool, []pb.WaitTxn, error) { + if err := ctx.Err(); err != nil { + return false, nil, err + } txn := s.activeTxnHolder.getActiveTxn(txnID, false, "") if txn == nil { return false, nil, nil @@ -31,7 +34,8 @@ func (s *service) GetWaitingList( v := txn.toWaitTxn(s.serviceID, false) if v.CreatedOn == s.serviceID { values := make([]pb.WaitTxn, 0, 1) - txn.fetchWhoWaitingMe( + _, err := txn.fetchWhoWaitingMe( + ctx, s.serviceID, txnID, func(w pb.WaitTxn, waiterAddress string) bool { @@ -39,13 +43,18 @@ func (s *service) GetWaitingList( values = append(values, w) return true }, - s.getLockTable) + func(_ context.Context, group uint32, table uint64) (lockTable, error) { + return s.getLockTable(ctx, group, table) + }) + if err != nil { + return false, nil, err + } return true, values, nil } - waitingList, err := s.getTxnWaitingListOnRemote(txnID, v.CreatedOn) + waitingList, err := s.getTxnWaitingListOnRemote(ctx, txnID, v.CreatedOn) if err != nil { - return false, nil, nil + return false, nil, err } return true, waitingList, nil } @@ -55,12 +64,20 @@ func (s *service) GetLockHolder( tableID uint64, row []byte, options pb.LockOptions) (pb.WaitTxn, bool, error) { - s.wait() + if err := ctx.Err(); err != nil { + return pb.WaitTxn{}, false, err + } + if err := s.wait(ctx); err != nil { + return pb.WaitTxn{}, false, err + } + if err := ctx.Err(); err != nil { + return pb.WaitTxn{}, false, err + } for { if err := ctx.Err(); err != nil { return pb.WaitTxn{}, false, err } - l, err := s.getLockTableWithCreate(options.Group, tableID, [][]byte{row}, options.Sharding) + l, err := s.getLockTableWithCreate(ctx, options.Group, tableID, [][]byte{row}, options.Sharding) if err != nil { if moerr.IsMoErrCode(err, moerr.ErrLockTableBindChanged) { continue @@ -139,7 +156,7 @@ func (s *service) ForceRefreshLockTableBinds( func (s *service) GetLockTableBind( group uint32, tableID uint64) (pb.LockTable, error) { - l, err := s.getLockTable(group, tableID) + l, err := s.getLockTable(context.Background(), group, tableID) if err != nil { return pb.LockTable{}, err } @@ -151,7 +168,8 @@ func (s *service) GetLockTableBind( func (s *service) GetLatestLockTableBind(bind pb.LockTable) (pb.LockTable, error) { requestAllocator := s.allocatorStateSnapshot() - newBind, allocator, err := getLockTableBind( + newBind, allocator, err := getLockTableBindWithContext( + context.Background(), s.remote.client, bind.Group, bind.Table, diff --git a/pkg/lockservice/service_observability_test.go b/pkg/lockservice/service_observability_test.go index eb9ee2330e9c8..84cc9bb0bcf76 100644 --- a/pkg/lockservice/service_observability_test.go +++ b/pkg/lockservice/service_observability_test.go @@ -442,9 +442,10 @@ func (l *getLockHolderTestTable) unlock( } func (l *getLockHolderTestTable) getLock( + _ context.Context, key []byte, txn pb.WaitTxn, - fn func(Lock)) { + fn func(Lock)) error { panic("unexpected getLock") } diff --git a/pkg/lockservice/service_remote.go b/pkg/lockservice/service_remote.go index 28f3093666472..0c5c700b6b1ef 100644 --- a/pkg/lockservice/service_remote.go +++ b/pkg/lockservice/service_remote.go @@ -275,12 +275,13 @@ func (s *service) handleRemoteLock( resp *pb.Response, cs morpc.ClientSession) { logFields := remoteLockResponseLogFields(req) - if !s.canLockOnServiceStatus(req.Lock.TxnID, req.Lock.Options, req.LockTable.Table, req.Lock.Rows) { + if !s.beginLockAdmission(req.Lock.TxnID, req.Lock.Options, req.LockTable.Table, req.Lock.Rows) { _ = writeResponseWithDeadline(s.logger, cancel, resp, moerr.NewRetryForCNRollingRestart(), cs, defaultRPCWriteTimeout, logFields) return } + defer s.endLockAdmission() - l, err := s.getLocalLockTable(req, resp) + l, err := s.getLocalLockTable(ctx, req, resp) if err != nil || l == nil { // means that the lockservice sending the lock request holds a stale @@ -378,12 +379,14 @@ func (s *service) handleForwardLock( resp *pb.Response, cs morpc.ClientSession) { logFields := remoteLockResponseLogFields(req) - if !s.canLockOnServiceStatus(req.Lock.TxnID, req.Lock.Options, req.LockTable.Table, req.Lock.Rows) { + if !s.beginLockAdmission(req.Lock.TxnID, req.Lock.Options, req.LockTable.Table, req.Lock.Rows) { _ = writeResponseWithDeadline(s.logger, cancel, resp, moerr.NewRetryForCNRollingRestart(), cs, defaultRPCWriteTimeout, logFields) return } + defer s.endLockAdmission() l, err := s.getLockTable( + ctx, req.LockTable.Group, req.LockTable.Table) if err != nil || @@ -555,7 +558,7 @@ func (s *service) handleRemoteGetLock( req *pb.Request, resp *pb.Response, cs morpc.ClientSession) { - l, err := s.getLocalLockTable(req, resp) + l, err := s.getLocalLockTable(ctx, req, resp) if err != nil || l == nil { // means that the lockservice sending the lock request holds a stale lock @@ -564,7 +567,8 @@ func (s *service) handleRemoteGetLock( return } - l.getLock( + err = l.getLock( + ctx, req.GetTxnLock.Row, pb.WaitTxn{TxnID: req.GetTxnLock.TxnID}, func(lock Lock) { @@ -591,7 +595,7 @@ func (s *service) handleRemoteGetLockHolder( req *pb.Request, resp *pb.Response, cs morpc.ClientSession) { - l, err := s.getLocalLockTable(req, resp) + l, err := s.getLocalLockTable(ctx, req, resp) if err != nil || l == nil { writeResponse(s.logger, cancel, resp, err, cs) return @@ -626,8 +630,9 @@ func (s *service) handleRemoteGetWaitingList( req *pb.Request, resp *pb.Response, cs morpc.ClientSession) { + txnID := bytes.Clone(req.GetWaitingList.Txn.TxnID) select { - case s.fetchWhoWaitingListC <- who{ctx: ctx, cancel: cancel, cs: cs, resp: resp, txnID: req.GetWaitingList.Txn.TxnID}: + case s.fetchWhoWaitingListC <- who{ctx: ctx, cancel: cancel, cs: cs, resp: resp, txnID: txnID}: return default: writeResponse(s.logger, cancel, resp, ErrDeadLockDetected, cs) @@ -651,7 +656,7 @@ func (s *service) handleKeepRemoteLock( req *pb.Request, resp *pb.Response, cs morpc.ClientSession) { - l, err := s.getLocalLockTable(req, resp) + l, err := s.getLocalLockTable(ctx, req, resp) if err != nil || l == nil { writeResponse(s.logger, cancel, resp, err, cs) @@ -664,9 +669,11 @@ func (s *service) handleKeepRemoteLock( } func (s *service) getLocalLockTable( + ctx context.Context, req *pb.Request, resp *pb.Response) (lockTable, error) { l, err := s.getLockTable( + ctx, req.LockTable.Group, req.LockTable.Table) if err != nil { @@ -675,6 +682,7 @@ func (s *service) getLocalLockTable( if l == nil { rows, sharding := lockTableLookupInputsFromRequest(req) l, err = s.getLockTableWithCreate( + ctx, req.LockTable.Group, req.LockTable.Table, rows, @@ -735,9 +743,10 @@ func lockTableLookupInputsFromRequest(req *pb.Request) ([][]byte, pb.Sharding) { } func (s *service) getTxnWaitingListOnRemote( + parent context.Context, txnID []byte, createdOn string) ([]pb.WaitTxn, error) { - ctx, cancel := context.WithTimeoutCause(context.Background(), defaultRPCTimeout, moerr.CauseGetTxnWaitingListOnRemote) + ctx, cancel := context.WithTimeoutCause(parent, defaultRPCTimeout, moerr.CauseGetTxnWaitingListOnRemote) defer cancel() req := acquireRequest() @@ -899,24 +908,6 @@ type allocatorState struct { version uint64 } -func getLockTableBind( - c Client, - group uint32, - tableID uint64, - originTableID uint64, - serviceID string, - sharding pb.Sharding) (pb.LockTable, allocatorState, error) { - return getLockTableBindWithContext( - context.Background(), - c, - group, - tableID, - originTableID, - serviceID, - sharding, - ) -} - func getLockTableBindWithContext( parent context.Context, c Client, @@ -960,6 +951,9 @@ type who struct { func (s *service) handleFetchWhoWaitingMe(ctx context.Context) { for { + if ctx.Err() != nil { + return + } select { case <-ctx.Done(): return @@ -972,7 +966,10 @@ func (s *service) handleFetchWhoWaitingMe(ctx context.Context) { writeResponse(s.logger, w.cancel, w.resp, nil, w.cs) continue } - txn.fetchWhoWaitingMe( + fetchCtx, fetchCancel := context.WithCancel(w.ctx) + stopServiceCancel := context.AfterFunc(ctx, fetchCancel) + _, fetchErr := txn.fetchWhoWaitingMe( + fetchCtx, s.serviceID, w.txnID, func(wt pb.WaitTxn, waiterAddress string) bool { @@ -980,8 +977,12 @@ func (s *service) handleFetchWhoWaitingMe(ctx context.Context) { w.resp.GetWaitingList.WaitingList = append(w.resp.GetWaitingList.WaitingList, wt) return true }, - s.getLockTable) - writeResponse(s.logger, w.cancel, w.resp, nil, w.cs) + func(ctx context.Context, group uint32, table uint64) (lockTable, error) { + return s.getLockTable(ctx, group, table) + }) + stopServiceCancel() + fetchCancel() + writeResponse(s.logger, w.cancel, w.resp, fetchErr, w.cs) } } } diff --git a/pkg/lockservice/service_remote_test.go b/pkg/lockservice/service_remote_test.go index 1c4a315dd444d..0e22d31dfdf01 100644 --- a/pkg/lockservice/service_remote_test.go +++ b/pkg/lockservice/service_remote_test.go @@ -137,7 +137,7 @@ func TestFetchWhoWaitingMeUsesActiveRemoteWaiterSnapshots(t *testing.T) { // test. Keep it bounded so an unexpected routing failure produces a // useful failure instead of consuming the package's 40-minute timeout. require.Eventually(t, func() bool { - lt, err := owner.getLockTable(0, tableID) + lt, err := owner.getLockTable(context.Background(), 0, tableID) if err != nil { return false } @@ -150,7 +150,7 @@ func TestFetchWhoWaitingMeUsesActiveRemoteWaiterSnapshots(t *testing.T) { lock, ok := local.mu.store.Get(row) return ok && lock.waiters.size() == 1 }, 10*time.Second, 10*time.Millisecond, "active remote waiter did not reach owner queue") - lt, err := owner.getLockTable(0, tableID) + lt, err := owner.getLockTable(context.Background(), 0, tableID) require.NoError(t, err) local := lt.(*localLockTable) logger := getLogger("") @@ -185,7 +185,8 @@ func TestFetchWhoWaitingMeUsesActiveRemoteWaiterSnapshots(t *testing.T) { txn := holderService.activeTxnHolder.getActiveTxn(holderTxn, false, "") require.NotNil(t, txn) var waitingTxnIDs [][]byte - require.True(t, txn.fetchWhoWaitingMe( + ok, err = txn.fetchWhoWaitingMe( + context.Background(), holderService.serviceID, holderTxn, func(waitTxn pb.WaitTxn, waiterAddress string) bool { @@ -193,8 +194,12 @@ func TestFetchWhoWaitingMeUsesActiveRemoteWaiterSnapshots(t *testing.T) { require.Equal(t, owner.serviceID, waiterAddress) return true }, - holderService.getLockTable, - )) + func(ctx context.Context, group uint32, table uint64) (lockTable, error) { + return holderService.getLockTable(ctx, group, table) + }, + ) + require.NoError(t, err) + require.True(t, ok) require.Equal(t, [][]byte{activeWaiterTxn}, waitingTxnIDs) releaseCtx, releaseCancel := context.WithTimeout(context.Background(), 10*time.Second) @@ -239,7 +244,7 @@ func TestWaitLocalWaitersIsBounded(t *testing.T) { txnID := []byte("holder") mustAddTestLock(t, ctx, s[0], tableID, txnID, [][]byte{row}, pb.Granularity_Row) - lt, err := s[0].getLockTable(0, tableID) + lt, err := s[0].getLockTable(ctx, 0, tableID) require.NoError(t, err) err = waitLocalWaitersWithTimeout(lt.(*localLockTable), row, 1, 20*time.Millisecond) require.EqualError(t, err, "internal error: timed out waiting for 1 local lock waiters, observed 0") @@ -267,7 +272,7 @@ func TestGetLocalLockTableUsesGetLockHolderLookupInputs(t *testing.T) { req.GetLockHolder.Sharding = pb.Sharding_ByRow resp := &pb.Response{} - lt, err := l.getLocalLockTable(req, resp) + lt, err := l.getLocalLockTable(context.Background(), req, resp) require.NoError(t, err) require.NotNil(t, lt) require.Equal(t, bind, lt.getBind()) @@ -1054,9 +1059,9 @@ func TestGetLockWithBindIsStable(t *testing.T) { table uint64) { txnID2 := []byte("txn2") - lt, err := l2.getLockTable(0, table) + lt, err := l2.getLockTable(context.Background(), 0, table) require.NoError(t, err) - lt.getLock(txnID2, pb.WaitTxn{TxnID: []byte{1}}, func(l Lock) {}) + _ = lt.getLock(context.Background(), txnID2, pb.WaitTxn{TxnID: []byte{1}}, func(l Lock) {}) checkBind( t, @@ -1151,9 +1156,9 @@ func TestGetLockWithBindTimeout(t *testing.T) { waitBindDisabled(t, alloc, l1.serviceID) txnID2 := []byte("txn2") - lt, err := l2.getLockTable(0, table) + lt, err := l2.getLockTable(context.Background(), 0, table) require.NoError(t, err) - lt.getLock(txnID2, pb.WaitTxn{TxnID: []byte{1}}, func(l Lock) {}) + _ = lt.getLock(context.Background(), txnID2, pb.WaitTxn{TxnID: []byte{1}}, func(l Lock) {}) // l2 get the bind l := l2.tableGroups.get(0, table) assert.Equal(t, l2.serviceID, l.getBind().ServiceID) @@ -1261,9 +1266,9 @@ func TestGetLockWithBindNotFound(t *testing.T) { }) txnID2 := []byte("txn2") - lt, err := l2.getLockTable(0, table) + lt, err := l2.getLockTable(context.Background(), 0, table) require.NoError(t, err) - lt.getLock(txnID2, pb.WaitTxn{TxnID: []byte{1}}, func(l Lock) {}) + _ = lt.getLock(context.Background(), txnID2, pb.WaitTxn{TxnID: []byte{1}}, func(l Lock) {}) checkBind( t, @@ -1386,7 +1391,7 @@ func TestIssue14346(t *testing.T) { case <-ctx.Done(): t.Fatal("timeout waiting for bind removal on s2") default: - v, err := s2.getLockTable(0, table) + v, err := s2.getLockTable(context.Background(), 0, table) require.NoError(t, err) if v == nil { return @@ -1425,14 +1430,14 @@ func runBindChangedTests( // l2 get the table1's bind mustAddTestLock(t, ctx, l2, table1, txnID2, [][]byte{{2}}, pb.Granularity_Row) - v, err := l2.getLockTable(0, table1) + v, err := l2.getLockTable(context.Background(), 0, table1) require.NoError(t, err) require.Equal(t, l1.serviceID, v.getBind().ServiceID) if makeBindChanged { // stop l1 keep lock bind skip.Store(true) - lt, err := l1.getLockTable(0, table1) + lt, err := l1.getLockTable(context.Background(), 0, table1) require.NoError(t, err) old := lt.getBind() waitBindDisabled(t, alloc, l1.serviceID) @@ -1476,7 +1481,7 @@ func waitBindChanged( old pb.LockTable, l *service) { for { - lt, err := l.getLockTableWithCreate(0, old.Table, nil, pb.Sharding_None) + lt, err := l.getLockTableWithCreate(context.Background(), 0, old.Table, nil, pb.Sharding_None) require.NoError(t, err) new := lt.getBind() if new.Changed(old) { diff --git a/pkg/lockservice/service_test.go b/pkg/lockservice/service_test.go index 755eb16f2713d..d88638b96ce53 100644 --- a/pkg/lockservice/service_test.go +++ b/pkg/lockservice/service_test.go @@ -66,7 +66,7 @@ func getRunner(remote bool) func(t *testing.T, table uint64, fn func(context.Con require.NoError(t, err, err) require.NoError(t, s1.Unlock(ctx, txn1, timestamp.Timestamp{})) - lt, err := s1.getLockTable(0, table) + lt, err := s1.getLockTable(context.Background(), 0, table) require.NoError(t, err) require.Equal(t, table, lt.getBind().Table) require.Equal(t, table, lt.getBind().OriginTable) @@ -1682,10 +1682,6 @@ func TestIssue3654(t *testing.T) { l1 := s[0] l2 := s[1] - ctx, cancel := context.WithTimeout( - context.Background(), - time.Nanosecond) - defer cancel() option := pb.LockOptions{ Granularity: pb.Granularity_Row, Mode: pb.LockMode_Exclusive, @@ -1694,13 +1690,18 @@ func TestIssue3654(t *testing.T) { } _, err := l1.Lock( - ctx, + context.Background(), 0, [][]byte{{1}}, []byte("txn1"), option) require.NoError(t, err) + ctx, cancel := context.WithTimeout( + context.Background(), + time.Nanosecond) + defer cancel() + _, err = l2.Lock( ctx, 0, @@ -2634,7 +2635,7 @@ func TestLockResultWithNoConflict(t *testing.T) { require.NoError(t, err) assert.False(t, res.Timestamp.IsEmpty()) - lb, err := l.getLockTable(0, 0) + lb, err := l.getLockTable(context.Background(), 0, 0) require.NoError(t, err) assert.Equal(t, lb.getBind(), res.LockedOn) }, @@ -3011,7 +3012,8 @@ func TestLeaveGetBindInRollingRestartCN(t *testing.T) { } } // get bind - _, _, err = getLockTableBind( + _, _, err = getLockTableBindWithContext( + ctx, l.remote.client, 0, 0, @@ -3503,7 +3505,7 @@ func TestGetBindPurgesStaleBindWhenAllocatorIDChangesWithRegressedVersion(t *tes restartedAllocatorID := alloc.allocatorID alloc.mu.Unlock() - _, err = l1.getLockTableWithCreate(0, freshTable, newTestRows(2), pb.Sharding_None) + _, err = l1.getLockTableWithCreate(context.Background(), 0, freshTable, newTestRows(2), pb.Sharding_None) require.NoError(t, err) require.Nil(t, l1.tableGroups.get(0, staleTable)) require.Equal(t, restartedVersion, l1.lastAllocatorVersion) @@ -3784,6 +3786,7 @@ func TestAllocatorPublishRejectsStaleBindAfterNewAllocatorObserved(t *testing.T) require.Nil(t, l1.tableGroups.get(0, staleTable)) lt, err := l1.publishLockTableBindFromAllocator( + context.Background(), "allocator-publish-race-old", staleBind.Group, staleBind.Table, @@ -3829,6 +3832,7 @@ func TestAllocatorPublishRejectsOverwriteAfterConcurrentBindChanged(t *testing.T l1.handleBindChanged(freshBind) lt, err := l1.publishLockTableBindFromAllocator( + context.Background(), "allocator-publish-current-race", delayedBind.Group, delayedBind.Table, @@ -4927,14 +4931,14 @@ func TestMultiGroupWithSameTableID(t *testing.T) { // txn1 get lock _, err := s.Lock(ctx, table, rows, txn1, option1) require.NoError(t, err) - lt1, err := s.getLockTable(g1, table) + lt1, err := s.getLockTable(context.Background(), g1, table) assert.NoError(t, err) checkLock(t, lt1.(*localLockTable), rows[0], [][]byte{txn1}, nil, nil) // txn2 get lock, shared _, err = s.Lock(ctx, table, rows, txn2, option2) require.NoError(t, err) - lt2, err := s.getLockTable(g2, table) + lt2, err := s.getLockTable(context.Background(), g2, table) assert.NoError(t, err) checkLock(t, lt2.(*localLockTable), rows[0], [][]byte{txn2}, nil, nil) @@ -5073,7 +5077,7 @@ func TestIssue2128(t *testing.T) { option) require.NoError(t, err) - lb, err := l.getLockTable(0, 0) + lb, err := l.getLockTable(context.Background(), 0, 0) require.NoError(t, err) b := lb.getBind() b.ServiceID = "1705661824807004000s3" @@ -5234,7 +5238,7 @@ func TestLeakWaiterForErr(t *testing.T) { txn3 := []byte("rt3") txn4 := []byte("rt4") - ll, err := l1.getLockTableWithCreate(0, tableID, nil, pb.Sharding_None) + ll, err := l1.getLockTableWithCreate(context.Background(), 0, tableID, nil, pb.Sharding_None) require.NoError(t, err) lt := ll.(*localLockTable) lt.options.afterWait = func(c *lockContext) func() { @@ -5313,7 +5317,7 @@ func TestIssue14008(t *testing.T) { wg.Add(1) go func() { defer wg.Done() - _, err := s1.getLockTableWithCreate(0, 10, nil, pb.Sharding_None) + _, err := s1.getLockTableWithCreate(context.Background(), 0, 10, nil, pb.Sharding_None) require.Error(t, err) }() } @@ -5462,7 +5466,7 @@ func TestLockWaitTimeoutSucceedsWhenHolderReleases(t *testing.T) { }() hasWaiter := func() bool { - v, err := l.getLockTable(0, 0) + v, err := l.getLockTable(context.Background(), 0, 0) require.NoError(t, err) lt := v.(*localLockTable) lt.mu.Lock() @@ -5857,3 +5861,123 @@ func mustAddTestLock(t *testing.T, lock, granularity) } + +type doneObservedContext struct { + context.Context + doneObserved chan struct{} + once sync.Once +} + +func (c *doneObservedContext) Done() <-chan struct{} { + c.once.Do(func() { close(c.doneObserved) }) + return c.Context.Done() +} + +func TestLockDoesNotAcquireAfterContextAlreadyCanceled(t *testing.T) { + runLockServiceTests(t, []string{"s1"}, + func(_ *lockTableAllocator, services []*service) { + s := services[0] + table := uint64(25790) + rows := newTestRows(1) + + warmupTxn := []byte("warm-bind") + _, err := s.Lock(context.Background(), table, rows, warmupTxn, + newTestRowExclusiveOptions()) + require.NoError(t, err) + require.NoError(t, s.Unlock(context.Background(), warmupTxn, + timestamp.Timestamp{})) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + canceledTxn := []byte("already-canceled") + _, lockErr := s.Lock(ctx, table, rows, canceledTxn, + newTestRowExclusiveOptions()) + if lockErr == nil { + require.NoError(t, s.Unlock(context.Background(), canceledTxn, + timestamp.Timestamp{})) + } + require.ErrorIs(t, lockErr, context.Canceled) + require.Nil(t, s.activeTxnHolder.getActiveTxn(canceledTxn, false, "")) + }) +} + +func TestLockReturnsWhenCanceledDuringBindAllocationWait(t *testing.T) { + runLockServiceTests(t, []string{"s1"}, + func(_ *lockTableAllocator, services []*service) { + s := services[0] + group := uint32(0) + table := uint64(257902) + waitC := make(chan struct{}) + s.mu.Lock() + s.mu.allocating[group] = map[uint64]chan struct{}{table: waitC} + s.mu.Unlock() + defer func() { + s.mu.Lock() + delete(s.mu.allocating[group], table) + s.mu.Unlock() + close(waitC) + }() + + baseCtx, cancel := context.WithCancel(context.Background()) + ctx := &doneObservedContext{ + Context: baseCtx, + doneObserved: make(chan struct{}), + } + txnID := []byte("canceled-bind-wait") + done := make(chan error, 1) + go func() { + _, err := s.Lock(ctx, table, newTestRows(1), txnID, + newTestRowExclusiveOptions()) + done <- err + }() + + <-ctx.doneObserved + s.checkCanMoveGroupTables() + require.True(t, s.isStatus(pb.Status_ServiceLockEnable), + "drain advanced while a lock admission was waiting for its bind") + cancel() + var lockErr error + select { + case lockErr = <-done: + case <-time.After(time.Second): + t.Fatal("Lock did not return after bind-allocation wait was canceled") + } + + require.ErrorIs(t, lockErr, context.Canceled) + require.Nil(t, s.activeTxnHolder.getActiveTxn(txnID, false, "")) + s.mu.RLock() + require.Zero(t, s.mu.lockAdmissions) + s.mu.RUnlock() + s.checkCanMoveGroupTables() + require.True(t, s.isStatus(pb.Status_ServiceLockWaiting)) + }) +} + +func TestLockReturnsWhenServiceReadinessWaitIsCanceled(t *testing.T) { + started := make(chan struct{}) + var once sync.Once + runLockServiceTests(t, []string{"s1"}, + func(_ *lockTableAllocator, services []*service) { + s := services[0] + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { + _, err := s.Lock(ctx, 257903, newTestRows(1), []byte("readiness-wait"), + newTestRowExclusiveOptions()) + done <- err + }() + <-started + cancel() + select { + case err := <-done: + require.ErrorIs(t, err, context.Canceled) + case <-time.After(time.Second): + t.Fatal("Lock did not return after the service-readiness wait was canceled") + } + }, + WithWait(func(ctx context.Context) error { + once.Do(func() { close(started) }) + <-ctx.Done() + return ctx.Err() + })) +} diff --git a/pkg/lockservice/test_helper.go b/pkg/lockservice/test_helper.go index 31f58619cf229..80754ae8075e0 100644 --- a/pkg/lockservice/test_helper.go +++ b/pkg/lockservice/test_helper.go @@ -15,6 +15,7 @@ package lockservice import ( + "context" "fmt" "os" "time" @@ -110,7 +111,7 @@ func WaitWaiters( key []byte, waitersCount int) error { s := ls.(*service) - v, err := s.getLockTable(group, table) + v, err := s.getLockTable(context.Background(), group, table) if err != nil { return err } diff --git a/pkg/lockservice/txn.go b/pkg/lockservice/txn.go index 6a616f65274f3..ed9049c59baf7 100644 --- a/pkg/lockservice/txn.go +++ b/pkg/lockservice/txn.go @@ -253,6 +253,9 @@ func (txn *activeTxn) closeWithContextInternal( for table, cs := range h.tableKeys { l, err := lockTableFunc(group, table) if err != nil { + if ctxErr := ctx.Err(); ctxErr != nil { + return ctxErr + } // if a remote transaction, then the corresponding locktable should be local // and cannot return an error. // @@ -478,15 +481,19 @@ func (txn *activeTxn) incLockTableRef(m map[uint32]map[uint64]uint64, serviceID // ============================================================================================================================ func (txn *activeTxn) fetchWhoWaitingMe( + ctx context.Context, serviceID string, txnID []byte, waiters func(pb.WaitTxn, string) bool, - lockTableFunc func(uint32, uint64) (lockTable, error)) bool { + lockTableFunc func(context.Context, uint32, uint64) (lockTable, error)) (bool, error) { + if err := ctx.Err(); err != nil { + return false, err + } txn.RLock() // txn already closed if !bytes.Equal(txn.txnID, txnID) { txn.RUnlock() - return true + return true, nil } // if this is a remote transaction, meaning that all the information is in the // remote, we need to execute the logic. @@ -516,14 +523,17 @@ func (txn *activeTxn) fetchWhoWaitingMe( }() for idx, table := range tables { - l, err := lockTableFunc(groups[idx], table) + if err := ctx.Err(); err != nil { + return false, err + } + l, err := lockTableFunc(ctx, groups[idx], table) if err != nil { // if a remote transaction, then the corresponding locktable should be local // and cannot return an error. // // or a local transaction holds a lock on remote lock table, but can not get // the remote LockTable, it is a bug. - panic(err) + return false, err } if l == nil { continue @@ -531,9 +541,15 @@ func (txn *activeTxn) fetchWhoWaitingMe( locks := lockKeys[idx] hasDeadLock := false + var fetchErr error waiterAddress := l.getBind().ServiceID locks.iter(func(lockKey []byte) bool { - l.getLock( + if err := ctx.Err(); err != nil { + fetchErr = err + return false + } + if err := l.getLock( + ctx, lockKey, wt, func(lock Lock) { @@ -548,15 +564,24 @@ func (txn *activeTxn) fetchWhoWaitingMe( hasDeadLock = !waiters(w.txn, waiterAddress) return !hasDeadLock }) - }) + }); err != nil { + fetchErr = err + return false + } return !hasDeadLock }) + if fetchErr != nil { + return false, fetchErr + } + if err := ctx.Err(); err != nil { + return false, err + } if hasDeadLock { - return false + return false, nil } } - return true + return true, nil } func (txn *activeTxn) toWaitTxn(serviceID string, locked bool) pb.WaitTxn { diff --git a/pkg/lockservice/txn_test.go b/pkg/lockservice/txn_test.go index a25b097ab55ec..d6dfcef97274f 100644 --- a/pkg/lockservice/txn_test.go +++ b/pkg/lockservice/txn_test.go @@ -68,7 +68,7 @@ func (l *retryableUnlockTestTable) unlockWithContext( return nil } -func (l *retryableUnlockTestTable) getLock([]byte, pb.WaitTxn, func(Lock)) { +func (l *retryableUnlockTestTable) getLock(context.Context, []byte, pb.WaitTxn, func(Lock)) error { panic("unexpected getLock") } @@ -205,7 +205,8 @@ func TestFetchWhoWaitingMeSkipsInactiveWaiters(t *testing.T) { }) var waitingTxnIDs [][]byte - ok := txn.fetchWhoWaitingMe( + ok, err := txn.fetchWhoWaitingMe( + context.Background(), "origin", holderID, func(waitTxn pb.WaitTxn, waiterAddress string) bool { @@ -213,13 +214,14 @@ func TestFetchWhoWaitingMeSkipsInactiveWaiters(t *testing.T) { assert.Equal(t, bind.ServiceID, waiterAddress) return true }, - func(group uint32, table uint64) (lockTable, error) { + func(_ context.Context, group uint32, table uint64) (lockTable, error) { assert.Equal(t, bind.Group, group) assert.Equal(t, bind.Table, table) return lt, nil }, ) + assert.NoError(t, err) assert.True(t, ok) assert.Equal(t, [][]byte{[]byte("blocking")}, waitingTxnIDs) }) @@ -309,3 +311,29 @@ func TestCloseWithoutFreeWithContextRetriesOnlyFailedTables(t *testing.T) { require.Empty(t, txn.lockHolders) }) } + +func TestCloseWithoutFreeWithContextReturnsCanceledLookup(t *testing.T) { + reuse.RunReuseTests(func() { + id := []byte("canceled-lookup") + txn := newActiveTxn(id, string(id), newFixedSlicePool(2), "") + defer reuse.Free(txn, nil) + + bind := pb.LockTable{Group: 0, Table: 1} + require.NoError(t, txn.lockAdded(0, bind, [][]byte{[]byte("k1")}, getLogger(""))) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + err := txn.closeWithoutFreeWithContext( + ctx, + id, + timestamp.Timestamp{}, + func(uint32, uint64) (lockTable, error) { + return nil, ctx.Err() + }, + getLogger(""), + ) + require.ErrorIs(t, err, context.Canceled) + require.Contains(t, txn.getHoldLocksLocked(0).tableKeys, bind.Table, + "canceled cleanup must retain the table for a later retry") + }) +} diff --git a/pkg/lockservice/types.go b/pkg/lockservice/types.go index 7ab32945df4f3..d268a6653ad36 100644 --- a/pkg/lockservice/types.go +++ b/pkg/lockservice/types.go @@ -195,7 +195,7 @@ type lockTable interface { // Unlock release a set of locks, if txn was committed, commitTS is not empty unlock(txn *activeTxn, ls *cowSlice, commitTS timestamp.Timestamp, mutations ...pb.ExtraMutation) // getLock get a lock - getLock(key []byte, txn pb.WaitTxn, fn func(Lock)) + getLock(ctx context.Context, key []byte, txn pb.WaitTxn, fn func(Lock)) error // getLockHolder returns the current holder if the lock is actively held. getLockHolder(ctx context.Context, key []byte) (pb.WaitTxn, bool, error) // getBind returns lock table binding diff --git a/pkg/lockservice/waiter.go b/pkg/lockservice/waiter.go index 65a4d2ae798dd..7520305a12f59 100644 --- a/pkg/lockservice/waiter.go +++ b/pkg/lockservice/waiter.go @@ -54,6 +54,7 @@ func acquireWaiter( panic("BUG: invalid ref count") } w.beforeSwapStatusAdjustFunc = func() {} + w.beforeWaitNotificationReceiveFunc = func() {} return w } @@ -114,7 +115,8 @@ type waiter struct { isRemoteSnapshot bool // just used for testing - beforeSwapStatusAdjustFunc func() + beforeSwapStatusAdjustFunc func() + beforeWaitNotificationReceiveFunc func() } // String implement Stringer @@ -185,16 +187,11 @@ func (w *waiter) casStatus( } func (w *waiter) mustRecvNotification( - ctx context.Context, logger *log.MOLogger, ) notifyValue { - select { - case v := <-w.c: - logWaiterGetNotify(logger, w, v) - return v - case <-ctx.Done(): - return notifyValue{err: ctx.Err()} - } + v := <-w.c + logWaiterGetNotify(logger, w, v) + return v } func (w *waiter) mustSendNotification( @@ -253,14 +250,17 @@ func (w *waiter) wait( w.beforeSwapStatusAdjustFunc() - // context is timeout, and status not changed, no concurrent happen - if w.casStatus(status, completed, logger) { + // Cancellation only owns the waiter if it can complete a still-blocking + // wait. Once a notifier has published notified, it owns completion and its + // channel value must be consumed even though ctx is already done. + if w.casStatus(blocking, completed, logger) { return notifyValue{err: ctx.Err()} } - // notify and timeout are concurrently issued, we use real result to replace - // timeout error + // Notification and cancellation raced; notification won the status claim. + w.beforeWaitNotificationReceiveFunc() + v := w.mustRecvNotification(logger) w.setStatus(completed) - return w.mustRecvNotification(ctx, logger) + return v } func (w *waiter) disableNotify() { @@ -330,6 +330,8 @@ func (w *waiter) reset() { w.lockWaitGranularity = pb.Granularity_Row w.lockWaitMode = pb.LockMode_Exclusive w.isRemoteSnapshot = false + w.beforeSwapStatusAdjustFunc = func() {} + w.beforeWaitNotificationReceiveFunc = func() {} w.stopLockWaitTimer() } diff --git a/pkg/lockservice/waiter_test.go b/pkg/lockservice/waiter_test.go index ee71cb9fe63f0..2def5b47d60a8 100644 --- a/pkg/lockservice/waiter_test.go +++ b/pkg/lockservice/waiter_test.go @@ -16,6 +16,7 @@ package lockservice import ( "context" + "errors" "testing" "time" @@ -80,6 +81,35 @@ func TestWaitAndNotifyConcurrent(t *testing.T) { } +func TestWaitCancellationAfterNotifyClaimConsumesNotification(t *testing.T) { + reuse.RunReuseTests(func() { + w := acquireWaiter(pb.WaitTxn{TxnID: []byte("w")}, "", nil) + defer w.close("", nil) + w.setStatus(notified) + + beforeReceive := make(chan struct{}) + allowReceive := make(chan struct{}) + w.beforeWaitNotificationReceiveFunc = func() { + close(beforeReceive) + <-allowReceive + } + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + expected := errors.New("notification won") + done := make(chan notifyValue, 1) + go func() { + done <- w.wait(ctx, getLogger("")) + }() + + <-beforeReceive + close(allowReceive) + w.c <- notifyValue{err: expected} + require.ErrorIs(t, (<-done).err, expected) + require.Empty(t, w.c) + }) +} + func TestWaitMultiTimes(t *testing.T) { reuse.RunReuseTests(func() { w := acquireWaiter(pb.WaitTxn{TxnID: []byte("w")}, "", nil) From 1c7910fa2368d3037a294081d61188434f48f2ca Mon Sep 17 00:00:00 2001 From: aptend Date: Mon, 20 Jul 2026 18:05:54 +0800 Subject: [PATCH 2/4] fix(lockservice): remove obsolete retry helper --- pkg/lockservice/lock_table_remote.go | 6 ------ 1 file changed, 6 deletions(-) diff --git a/pkg/lockservice/lock_table_remote.go b/pkg/lockservice/lock_table_remote.go index 63ec90bae239b..8ed4a85bbe9bf 100644 --- a/pkg/lockservice/lock_table_remote.go +++ b/pkg/lockservice/lock_table_remote.go @@ -339,12 +339,6 @@ func (l *remoteLockTable) getLockHolder(ctx context.Context, key []byte) (pb.Wai } } -func waitRemoteRetryBackoff(backoff time.Duration) { - if backoff > 0 { - time.Sleep(backoff) - } -} - func waitRemoteRetryBackoffWithContext(ctx context.Context, backoff time.Duration) error { if backoff <= 0 { return ctx.Err() From 2cbd07dddd02d97cfcccb5eab1f429c47b1f412b Mon Sep 17 00:00:00 2001 From: aptend Date: Mon, 20 Jul 2026 23:14:06 +0800 Subject: [PATCH 3/4] fix(lockservice): linearize rolling drain admissions --- pkg/lockservice/service.go | 191 +++++++++++++++++++--------- pkg/lockservice/service_remote.go | 44 ++----- pkg/lockservice/service_test.go | 203 +++++++++++++++++++++++++++++- pkg/lockservice/txn.go | 52 ++++++-- pkg/lockservice/txn_test.go | 5 +- 5 files changed, 390 insertions(+), 105 deletions(-) diff --git a/pkg/lockservice/service.go b/pkg/lockservice/service.go index d074f21fe4fdb..2de5d544c53a3 100644 --- a/pkg/lockservice/service.go +++ b/pkg/lockservice/service.go @@ -77,12 +77,16 @@ type service struct { mu struct { sync.RWMutex - restartTime timestamp.Timestamp - status pb.Status - lockAdmissions uint64 - groupTables [][]pb.LockTable - lockTableRef map[uint32]map[uint64]uint64 - allocating map[uint32]map[uint64]chan struct{} + restartTime timestamp.Timestamp + status pb.Status + lockAdmissions uint64 + preDrainAdmissions uint64 + txnClosures uint64 + drainSnapshotReady bool + groupTables [][]pb.LockTable + lockTableRef map[uint32]map[uint64]uint64 + pendingRefReleases []pb.LockTable + allocating map[uint32]map[uint64]chan struct{} } option struct { @@ -162,10 +166,11 @@ func (s *service) Lock( return pb.Result{}, err } - if !s.beginLockAdmission(txnID, options, tableID, rows) { + admitted, preDrain := s.beginLockAdmission(txnID, options, tableID, rows) + if !admitted { return pb.Result{}, moerr.NewNewTxnInCNRollingRestart() } - defer s.endLockAdmission() + defer s.endLockAdmission(preDrain) v2.TxnLockTotalCounter.Inc() options.Validate(rows) @@ -233,19 +238,11 @@ func (s *service) Lock( s.bindChangeMu.RUnlock() return pb.Result{}, ErrLockTableBindChanged } - h := txn.getHoldLocksLocked(bind.Group) - _, hasBind := h.tableBinds[bind.Table] - txn.lockTableBindTouched(bind) + if txn.lockTableBindTouched(bind) && bind.ServiceID == s.serviceID { + s.incRef(bind.Group, bind.Table) + } s.bindChangeMu.RUnlock() defer txn.Unlock() - defer func() { - if s.isStatus(pb.Status_ServiceLockEnable) || - err != nil || - hasBind { - return - } - s.incRef(bind.Group, bind.Table) - }() var result pb.Result l.lock( @@ -316,6 +313,7 @@ func (s *service) unlockUnknownCommit( } defer logUnlockTxn(s.logger, txn)() + binds := txn.lockTableBindsLocked() if err := txn.closeWithoutFreeWithContext( ctx, txnID, @@ -332,10 +330,8 @@ func (s *service) unlockUnknownCommit( if s.activeTxnHolder.deleteActiveTxn(txnID) != txn { return nil } - if !s.isStatus(pb.Status_ServiceLockEnable) { - s.reduceCanMoveGroupTables(txn) - s.tryCompleteDrain() - } + s.reduceCanMoveGroupTables(binds) + s.tryCompleteDrain() s.deadlockDetector.txnClosed(txnID) reuse.Free(txn, nil) return nil @@ -361,6 +357,9 @@ func (s *service) unlockWithContext( return err } + s.beginTxnClosure() + defer s.endTxnClosure() + txn := s.activeTxnHolder.deleteActiveTxn(txnID) if txn == nil { return nil @@ -372,21 +371,21 @@ func (s *service) unlockWithContext( return nil } - if !s.isStatus(pb.Status_ServiceLockEnable) { - s.reduceCanMoveGroupTables(txn) - s.tryCompleteDrain() - } - defer logUnlockTxn(s.logger, txn)() + binds := txn.lockTableBindsLocked() err := txn.closeWithContext(ctx, txnID, commitTS, func(group uint32, table uint64) (lockTable, error) { return s.getLockTable(ctx, group, table) }, s.logger, mutations...) + if err != nil { + return err + } + s.reduceCanMoveGroupTables(binds) // The deadlock detector will hold the deadlocked transaction that is aborted // to avoid the situation where the deadlock detection is interfered with by // the abort transaction. When a transaction is unlocked, the deadlock detector // needs to be notified to release memory. s.deadlockDetector.txnClosed(txnID) - return err + return nil } func (s *service) IsOrphanTxn( @@ -428,41 +427,82 @@ func (s *service) Resume() error { return err } -func (s *service) reduceCanMoveGroupTables(txn *activeTxn) { +func (s *service) reduceCanMoveGroupTables(binds []pb.LockTable) { s.mu.Lock() defer s.mu.Unlock() if len(s.mu.lockTableRef) == 0 { return } - var res []pb.LockTable - - for group, h := range txn.lockHolders { - for table, bind := range h.tableBinds { - if bind.ServiceID == s.serviceID { - if _, ok := s.mu.lockTableRef[group][table]; ok { - s.mu.lockTableRef[group][table]-- - if s.mu.lockTableRef[group][table] == 0 { - delete(s.mu.lockTableRef[group], table) - res = append(res, bind) - } - } - } + for _, bind := range binds { + if bind.ServiceID != s.serviceID { + continue + } + if s.mu.lockAdmissions != 0 { + s.mu.pendingRefReleases = append(s.mu.pendingRefReleases, bind) + continue } + s.releaseBindRefLocked(bind.Group, bind.Table, bind, s.mu.drainSnapshotReady) } - if len(res) > 0 { - s.mu.groupTables = append(s.mu.groupTables, res) +} + +func (s *service) releaseBindRefLocked( + group uint32, + table uint64, + bind pb.LockTable, + addMovable bool, +) { + if _, ok := s.mu.lockTableRef[group][table]; !ok { + return + } + s.mu.lockTableRef[group][table]-- + if s.mu.lockTableRef[group][table] != 0 { + return + } + delete(s.mu.lockTableRef[group], table) + if addMovable { + s.mu.groupTables = append(s.mu.groupTables, []pb.LockTable{bind}) } } +func (s *service) applyPendingRefReleasesLocked() { + if s.mu.lockAdmissions != 0 || len(s.mu.pendingRefReleases) == 0 { + return + } + for _, bind := range s.mu.pendingRefReleases { + s.releaseBindRefLocked( + bind.Group, + bind.Table, + bind, + s.mu.drainSnapshotReady, + ) + } + s.mu.pendingRefReleases = s.mu.pendingRefReleases[:0] +} + func (s *service) checkCanMoveGroupTables() { s.mu.Lock() defer s.mu.Unlock() - if s.mu.status != pb.Status_ServiceLockEnable || s.mu.lockAdmissions > 0 { + if s.mu.status != pb.Status_ServiceLockEnable { + return + } + + oldStatus := s.mu.status + s.mu.restartTime, _ = s.clock.Now() + s.mu.status = pb.Status_ServiceLockWaiting + s.mu.preDrainAdmissions = s.mu.lockAdmissions + s.mu.drainSnapshotReady = false + logStatusChange(s.logger, oldStatus, s.mu.status) + s.prepareDrainSnapshotLocked() +} + +func (s *service) prepareDrainSnapshotLocked() { + if s.mu.status != pb.Status_ServiceLockWaiting || + s.mu.drainSnapshotReady || + s.mu.preDrainAdmissions != 0 { return } - s.activeTxnHolder.incLockTableRef(s.mu.lockTableRef, s.serviceID) var res []pb.LockTable s.tableGroups.iter(func(_ uint64, v lockTable) bool { bind := v.getBind() @@ -476,9 +516,7 @@ func (s *service) checkCanMoveGroupTables() { if len(res) > 0 { s.mu.groupTables = append(s.mu.groupTables, res) } - s.mu.restartTime, _ = s.clock.Now() - s.mu.status = pb.Status_ServiceLockWaiting - logStatusChange(s.logger, s.mu.status, pb.Status_ServiceLockWaiting) + s.mu.drainSnapshotReady = true } func (s *service) incRef(group uint32, table uint64) { @@ -502,13 +540,13 @@ func (s *service) canLockOnServiceStatusLocked( if opts.Sharding == pb.Sharding_ByRow { tableID = ShardingByRow(rows[0]) } - if s.activeTxnHolder.hasActiveTxn(txnID) { - return true - } if _, ok := s.mu.lockTableRef[opts.Group][tableID]; !ok { logCanLockOnService(s.logger, s.serviceID) return false } + if s.activeTxnHolder.hasActiveTxn(txnID) { + return true + } if s.activeTxnHolder.empty() { return false } @@ -523,35 +561,49 @@ func (s *service) beginLockAdmission( opts pb.LockOptions, tableID uint64, rows [][]byte, -) bool { +) (bool, bool) { s.mu.Lock() defer s.mu.Unlock() if !s.canLockOnServiceStatusLocked(txnID, opts, tableID, rows) { - return false + return false, false } + preDrain := s.mu.status == pb.Status_ServiceLockEnable s.mu.lockAdmissions++ - return true + return true, preDrain } -func (s *service) endLockAdmission() { +func (s *service) endLockAdmission(preDrain bool) { s.mu.Lock() defer s.mu.Unlock() if s.mu.lockAdmissions == 0 { panic("lock admission underflow") } s.mu.lockAdmissions-- + if preDrain && s.mu.status == pb.Status_ServiceLockWaiting { + if s.mu.preDrainAdmissions == 0 { + panic("pre-drain lock admission underflow") + } + s.mu.preDrainAdmissions-- + } + s.applyPendingRefReleasesLocked() + s.prepareDrainSnapshotLocked() s.tryCompleteDrainLocked() } func (s *service) tryCompleteDrain() { s.mu.Lock() defer s.mu.Unlock() + s.applyPendingRefReleasesLocked() + s.prepareDrainSnapshotLocked() s.tryCompleteDrainLocked() } func (s *service) tryCompleteDrainLocked() { if s.mu.status != pb.Status_ServiceLockWaiting || + !s.mu.drainSnapshotReady || s.mu.lockAdmissions != 0 || + s.mu.txnClosures != 0 || + s.hasLockTableRefsLocked() || !s.activeTxnHolder.empty() { return } @@ -559,6 +611,31 @@ func (s *service) tryCompleteDrainLocked() { s.mu.status = pb.Status_ServiceUnLockSucc } +func (s *service) hasLockTableRefsLocked() bool { + for _, refs := range s.mu.lockTableRef { + if len(refs) != 0 { + return true + } + } + return false +} + +func (s *service) beginTxnClosure() { + s.mu.Lock() + s.mu.txnClosures++ + s.mu.Unlock() +} + +func (s *service) endTxnClosure() { + s.mu.Lock() + defer s.mu.Unlock() + if s.mu.txnClosures == 0 { + panic("transaction closure underflow") + } + s.mu.txnClosures-- + s.tryCompleteDrainLocked() +} + func (s *service) validGroupTable(group uint32, tableID uint64) bool { s.mu.RLock() defer s.mu.RUnlock() diff --git a/pkg/lockservice/service_remote.go b/pkg/lockservice/service_remote.go index 0c5c700b6b1ef..e26bf459fc0c5 100644 --- a/pkg/lockservice/service_remote.go +++ b/pkg/lockservice/service_remote.go @@ -275,11 +275,12 @@ func (s *service) handleRemoteLock( resp *pb.Response, cs morpc.ClientSession) { logFields := remoteLockResponseLogFields(req) - if !s.beginLockAdmission(req.Lock.TxnID, req.Lock.Options, req.LockTable.Table, req.Lock.Rows) { + admitted, preDrain := s.beginLockAdmission(req.Lock.TxnID, req.Lock.Options, req.LockTable.Table, req.Lock.Rows) + if !admitted { _ = writeResponseWithDeadline(s.logger, cancel, resp, moerr.NewRetryForCNRollingRestart(), cs, defaultRPCWriteTimeout, logFields) return } - defer s.endLockAdmission() + defer s.endLockAdmission(preDrain) l, err := s.getLocalLockTable(ctx, req, resp) if err != nil || @@ -333,22 +334,12 @@ func (s *service) handleRemoteLock( return } - var lockErr error - // it needs to inc table bind ref when set restart cn - h := txn.getHoldLocksLocked(bind.Group) - _, hasBind := h.tableBinds[bind.Table] - txn.lockTableBindTouched(bind) + if txn.lockTableBindTouched(bind) && bind.ServiceID == s.serviceID { + s.incRef(bind.Group, bind.Table) + } txnID := append([]byte(nil), req.Lock.TxnID...) s.bindChangeMu.RUnlock() defer txn.Unlock() - defer func() { - if s.isStatus(pb.Status_ServiceLockEnable) || - lockErr != nil || - hasBind { - return - } - s.incRef(bind.Group, bind.Table) - }() l.lock( ctx, @@ -366,7 +357,6 @@ func (s *service) handleRemoteLock( err = e } } - lockErr = err resp.Lock.Result = result _ = writeResponseWithDeadline(s.logger, cancel, resp, err, cs, defaultRPCWriteTimeout, logFields) }) @@ -379,11 +369,12 @@ func (s *service) handleForwardLock( resp *pb.Response, cs morpc.ClientSession) { logFields := remoteLockResponseLogFields(req) - if !s.beginLockAdmission(req.Lock.TxnID, req.Lock.Options, req.LockTable.Table, req.Lock.Rows) { + admitted, preDrain := s.beginLockAdmission(req.Lock.TxnID, req.Lock.Options, req.LockTable.Table, req.Lock.Rows) + if !admitted { _ = writeResponseWithDeadline(s.logger, cancel, resp, moerr.NewRetryForCNRollingRestart(), cs, defaultRPCWriteTimeout, logFields) return } - defer s.endLockAdmission() + defer s.endLockAdmission(preDrain) l, err := s.getLockTable( ctx, @@ -435,22 +426,12 @@ func (s *service) handleForwardLock( return } - var lockErr error - // it needs to inc table bind ref when set restart cn - h := txn.getHoldLocksLocked(bind.Group) - _, hasBind := h.tableBinds[bind.Table] - txn.lockTableBindTouched(bind) + if txn.lockTableBindTouched(bind) && bind.ServiceID == s.serviceID { + s.incRef(bind.Group, bind.Table) + } txnID := append([]byte(nil), req.Lock.TxnID...) s.bindChangeMu.RUnlock() defer txn.Unlock() - defer func() { - if s.isStatus(pb.Status_ServiceLockEnable) || - lockErr != nil || - hasBind { - return - } - s.incRef(bind.Group, bind.Table) - }() l.lock( ctx, @@ -468,7 +449,6 @@ func (s *service) handleForwardLock( err = e } } - lockErr = err resp.Lock.Result = result _ = writeResponseWithDeadline(s.logger, cancel, resp, err, cs, defaultRPCWriteTimeout, logFields) }) diff --git a/pkg/lockservice/service_test.go b/pkg/lockservice/service_test.go index d88638b96ce53..e73ed0ed3a2ed 100644 --- a/pkg/lockservice/service_test.go +++ b/pkg/lockservice/service_test.go @@ -5868,6 +5868,24 @@ type doneObservedContext struct { once sync.Once } +type blockingUnlockTestTable struct { + retryableUnlockTestTable + started chan struct{} + release chan struct{} +} + +func (l *blockingUnlockTestTable) unlockWithContext( + context.Context, + *activeTxn, + *cowSlice, + timestamp.Timestamp, + ...pb.ExtraMutation, +) error { + close(l.started) + <-l.release + return nil +} + func (c *doneObservedContext) Done() <-chan struct{} { c.once.Do(func() { close(c.doneObserved) }) return c.Context.Done() @@ -5933,8 +5951,16 @@ func TestLockReturnsWhenCanceledDuringBindAllocationWait(t *testing.T) { <-ctx.doneObserved s.checkCanMoveGroupTables() - require.True(t, s.isStatus(pb.Status_ServiceLockEnable), - "drain advanced while a lock admission was waiting for its bind") + require.True(t, s.isStatus(pb.Status_ServiceLockWaiting), + "drain request must close the admission gate immediately") + s.mu.RLock() + require.False(t, s.mu.drainSnapshotReady, + "drain snapshot must wait for pre-gate admissions") + s.mu.RUnlock() + + _, err := s.Lock(context.Background(), table+1, newTestRows(1), + []byte("post-drain-attempt"), newTestRowExclusiveOptions()) + require.True(t, moerr.IsMoErrCode(err, moerr.ErrNewTxnInCNRollingRestart)) cancel() var lockErr error select { @@ -5947,9 +5973,182 @@ func TestLockReturnsWhenCanceledDuringBindAllocationWait(t *testing.T) { require.Nil(t, s.activeTxnHolder.getActiveTxn(txnID, false, "")) s.mu.RLock() require.Zero(t, s.mu.lockAdmissions) + require.True(t, s.mu.drainSnapshotReady) s.mu.RUnlock() + }) +} + +func TestDrainPinsAsyncRemoteWaiterIntent(t *testing.T) { + runLockServiceTests(t, []string{"s1", "s2"}, + func(_ *lockTableAllocator, services []*service) { + owner := services[0] + caller := services[1] + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + table := uint64(257904) + row := []byte{1} + options := newTestRowExclusiveOptions() + holderTxn := []byte("drain-holder") + waiterTxn := []byte("drain-remote-waiter") + _, err := owner.Lock(ctx, table, [][]byte{row}, holderTxn, options) + require.NoError(t, err) + + lockDone := make(chan error, 1) + go func() { + _, err := caller.Lock(ctx, table, [][]byte{row}, waiterTxn, options) + lockDone <- err + }() + waitWaiters(t, owner, table, row, 1) + + owner.checkCanMoveGroupTables() + require.True(t, owner.isStatus(pb.Status_ServiceLockWaiting)) + require.True(t, owner.validGroupTable(0, table), + "pending remote waiter intent must pin the table") + + require.NoError(t, owner.Unlock(ctx, holderTxn, timestamp.Timestamp{})) + require.True(t, owner.validGroupTable(0, table), + "table must remain pinned after the original holder unlocks") + select { + case err := <-lockDone: + require.NoError(t, err) + case <-ctx.Done(): + t.Fatal("remote waiter did not acquire after holder unlock") + } + require.NoError(t, caller.Unlock(ctx, waiterTxn, timestamp.Timestamp{})) + }) +} + +func TestCompletedTxnDoesNotLeaveDrainRef(t *testing.T) { + runLockServiceTests(t, []string{"s1"}, + func(_ *lockTableAllocator, services []*service) { + s := services[0] + table := uint64(257905) + txnID := []byte("completed-before-drain") + _, err := s.Lock(context.Background(), table, newTestRows(1), txnID, + newTestRowExclusiveOptions()) + require.NoError(t, err) + require.NoError(t, s.Unlock(context.Background(), txnID, timestamp.Timestamp{})) + + s.mu.RLock() + _, pinned := s.mu.lockTableRef[0][table] + s.mu.RUnlock() + require.False(t, pinned) + s.checkCanMoveGroupTables() + movable := s.topGroupTables() + require.Len(t, movable, 1) + require.Equal(t, table, movable[0].Table) + }) +} + +func TestPostGateAdmissionsDoNotBlockDrainSnapshot(t *testing.T) { + runLockServiceTests(t, []string{"s1"}, + func(_ *lockTableAllocator, services []*service) { + s := services[0] + table := uint64(257906) + txnID := []byte("pinned-active-txn") + options := newTestRowExclusiveOptions() + _, err := s.Lock(context.Background(), table, newTestRows(1), txnID, options) + require.NoError(t, err) + + admitted, preDrain := s.beginLockAdmission(txnID, options, table, newTestRows(1)) + require.True(t, admitted) + require.True(t, preDrain) + s.checkCanMoveGroupTables() + + for range 100 { + admitted, postDrain := s.beginLockAdmission(txnID, options, table, newTestRows(1)) + require.True(t, admitted) + require.False(t, postDrain) + s.endLockAdmission(postDrain) + } + s.mu.RLock() + require.False(t, s.mu.drainSnapshotReady) + require.Equal(t, uint64(1), s.mu.preDrainAdmissions) + s.mu.RUnlock() + + s.endLockAdmission(preDrain) + s.mu.RLock() + require.True(t, s.mu.drainSnapshotReady) + require.Zero(t, s.mu.preDrainAdmissions) + s.mu.RUnlock() + require.NoError(t, s.Unlock(context.Background(), txnID, timestamp.Timestamp{})) + }) +} + +func TestDrainWaitsForTxnCloseLinearization(t *testing.T) { + runLockServiceTests(t, []string{"s1"}, + func(_ *lockTableAllocator, services []*service) { + s := services[0] + table := uint64(257907) + bind := pb.LockTable{ + Group: 0, Table: table, OriginTable: table, + ServiceID: s.serviceID, Valid: true, Version: 1, + } + lt := &blockingUnlockTestTable{ + retryableUnlockTestTable: retryableUnlockTestTable{bind: bind}, + started: make(chan struct{}), + release: make(chan struct{}), + } + s.tableGroups.set(0, table, lt) + + txnID := []byte("blocking-drain-close") + txn := s.activeTxnHolder.getActiveTxn(txnID, true, "") + txn.Lock() + require.True(t, txn.lockTableBindTouched(bind)) + s.incRef(bind.Group, bind.Table) + require.NoError(t, txn.lockAdded(bind.Group, bind, [][]byte{{1}}, s.logger)) + txn.Unlock() + s.checkCanMoveGroupTables() + + done := make(chan error, 1) + go func() { + done <- s.unlockWithContext(context.Background(), txnID, timestamp.Timestamp{}) + }() + <-lt.started require.True(t, s.isStatus(pb.Status_ServiceLockWaiting)) + s.mu.RLock() + require.Equal(t, uint64(1), s.mu.txnClosures) + s.mu.RUnlock() + + close(lt.release) + require.NoError(t, <-done) + require.True(t, s.isStatus(pb.Status_ServiceUnLockSucc)) + }) +} + +func TestRetryableUnknownCommitCloseReleasesAllDrainRefs(t *testing.T) { + runLockServiceTests(t, []string{"s1"}, + func(_ *lockTableAllocator, services []*service) { + s := services[0] + txnID := []byte("retryable-drain-refs") + txn := s.activeTxnHolder.getActiveTxn(txnID, true, "") + tables := map[uint64]*retryableUnlockTestTable{ + 257908: {bind: pb.LockTable{Group: 0, Table: 257908, ServiceID: s.serviceID, Valid: true}}, + 257909: {bind: pb.LockTable{Group: 0, Table: 257909, ServiceID: s.serviceID, Valid: true}, failFirst: true}, + } + + txn.Lock() + for table, lt := range tables { + s.tableGroups.set(0, table, lt) + require.True(t, txn.lockTableBindTouched(lt.bind)) + s.incRef(lt.bind.Group, lt.bind.Table) + require.NoError(t, txn.lockAdded(0, lt.bind, [][]byte{{byte(table)}}, s.logger)) + } + txn.Unlock() + + err := s.unlockUnknownCommit(context.Background(), txnID, timestamp.Timestamp{}) + require.ErrorIs(t, err, context.DeadlineExceeded) + require.NotNil(t, s.activeTxnHolder.getActiveTxn(txnID, false, "")) + for table := range tables { + require.True(t, s.validGroupTable(0, table)) + } + + require.NoError(t, s.unlockUnknownCommit(context.Background(), txnID, timestamp.Timestamp{})) + for table := range tables { + require.False(t, s.validGroupTable(0, table)) + } }) } diff --git a/pkg/lockservice/txn.go b/pkg/lockservice/txn.go index ed9049c59baf7..9845fd4dcab9f 100644 --- a/pkg/lockservice/txn.go +++ b/pkg/lockservice/txn.go @@ -153,13 +153,41 @@ func (txn *activeTxn) lockAdded( return nil } -func (txn *activeTxn) lockTableBindTouched(bind pb.LockTable) { +func (txn *activeTxn) lockTableBindTouched(bind pb.LockTable) bool { h := txn.getHoldLocksLocked(bind.Group) - if _, ok := h.tableBindIntents[bind.Table]; !ok { - h.tableBindIntents[bind.Table] = bind + if _, ok := h.tableBindIntents[bind.Table]; ok { + return false + } + h.tableBindIntents[bind.Table] = bind + return true +} + +// iterLockTableBindsLocked visits every table touched by the transaction once. +// Acquired binds take precedence over intents because they are authoritative. +// The caller must hold txn's mutex. +func (txn *activeTxn) iterLockTableBindsLocked( + fn func(group uint32, table uint64, bind pb.LockTable), +) { + for group, h := range txn.lockHolders { + for table, bind := range h.tableBinds { + fn(group, table, bind) + } + for table, bind := range h.tableBindIntents { + if _, ok := h.tableBinds[table]; !ok { + fn(group, table, bind) + } + } } } +func (txn *activeTxn) lockTableBindsLocked() []pb.LockTable { + binds := make([]pb.LockTable, 0, len(txn.lockHolders)) + txn.iterLockTableBindsLocked(func(_ uint32, _ uint64, bind pb.LockTable) { + binds = append(binds, bind) + }) + return binds +} + func (txn *activeTxn) close( txnID []byte, commitTS timestamp.Timestamp, @@ -336,7 +364,9 @@ func (txn *activeTxn) removeClosedLockTable( } delete(h.tableKeys, table) delete(h.tableBinds, table) - delete(h.tableBindIntents, table) + // Keep the intent until the whole transaction closes. It owns the service + // drain reference even after this table was successfully released during a + // retryable, multi-table cleanup. cs.close() if len(h.tableKeys) == 0 && len(h.tableBinds) == 0 && len(h.tableBindIntents) == 0 { @@ -463,16 +493,14 @@ func (txn *activeTxn) isRemoteLocked() bool { func (txn *activeTxn) incLockTableRef(m map[uint32]map[uint64]uint64, serviceID string) { txn.RLock() defer txn.RUnlock() - for _, h := range txn.lockHolders { - for _, l := range h.tableBinds { - if serviceID == l.ServiceID { - if _, ok := m[l.Group]; !ok { - m[l.Group] = make(map[uint64]uint64, 1024) - } - m[l.Group][l.Table]++ + txn.iterLockTableBindsLocked(func(_ uint32, _ uint64, l pb.LockTable) { + if serviceID == l.ServiceID { + if _, ok := m[l.Group]; !ok { + m[l.Group] = make(map[uint64]uint64, 1024) } + m[l.Group][l.Table]++ } - } + }) } // ============================================================================================================================ diff --git a/pkg/lockservice/txn_test.go b/pkg/lockservice/txn_test.go index d6dfcef97274f..4eeef2015cee0 100644 --- a/pkg/lockservice/txn_test.go +++ b/pkg/lockservice/txn_test.go @@ -130,7 +130,8 @@ func TestLockTableBindTouchedTracksFenceIntentOnly(t *testing.T) { defer reuse.Free(txn, nil) bind := pb.LockTable{Group: 0, Table: 1, ServiceID: "s1", Version: 1} - txn.lockTableBindTouched(bind) + assert.True(t, txn.lockTableBindTouched(bind)) + assert.False(t, txn.lockTableBindTouched(bind)) h := txn.getHoldLocksLocked(bind.Group) assert.Empty(t, h.tableBinds) @@ -138,7 +139,7 @@ func TestLockTableBindTouchedTracksFenceIntentOnly(t *testing.T) { refs := make(map[uint32]map[uint64]uint64) txn.incLockTableRef(refs, bind.ServiceID) - assert.Empty(t, refs) + assert.Equal(t, uint64(1), refs[bind.Group][bind.Table]) changed := bind changed.Version++ From 9395ec2312229b9c74a899f1df5e642ba2ce842e Mon Sep 17 00:00:00 2001 From: aptend Date: Tue, 21 Jul 2026 10:32:46 +0800 Subject: [PATCH 4/4] fix(lockservice): bound drain reference accounting --- pkg/lockservice/service.go | 81 +++++++++++++++++++------------ pkg/lockservice/service_remote.go | 16 +++--- pkg/lockservice/service_test.go | 48 ++++++++++++++++-- 3 files changed, 104 insertions(+), 41 deletions(-) diff --git a/pkg/lockservice/service.go b/pkg/lockservice/service.go index 2de5d544c53a3..4704b621d37e4 100644 --- a/pkg/lockservice/service.go +++ b/pkg/lockservice/service.go @@ -85,7 +85,6 @@ type service struct { drainSnapshotReady bool groupTables [][]pb.LockTable lockTableRef map[uint32]map[uint64]uint64 - pendingRefReleases []pb.LockTable allocating map[uint32]map[uint64]chan struct{} } @@ -166,11 +165,11 @@ func (s *service) Lock( return pb.Result{}, err } - admitted, preDrain := s.beginLockAdmission(txnID, options, tableID, rows) + admission, admitted := s.beginLockAdmission(txnID, options, tableID, rows) if !admitted { return pb.Result{}, moerr.NewNewTxnInCNRollingRestart() } - defer s.endLockAdmission(preDrain) + defer func() { s.endLockAdmission(admission) }() v2.TxnLockTotalCounter.Inc() options.Validate(rows) @@ -238,7 +237,9 @@ func (s *service) Lock( s.bindChangeMu.RUnlock() return pb.Result{}, ErrLockTableBindChanged } - if txn.lockTableBindTouched(bind) && bind.ServiceID == s.serviceID { + if txn.lockTableBindTouched(bind) && + bind.ServiceID == s.serviceID && + !admission.consume(bind) { s.incRef(bind.Group, bind.Table) } s.bindChangeMu.RUnlock() @@ -438,10 +439,6 @@ func (s *service) reduceCanMoveGroupTables(binds []pb.LockTable) { if bind.ServiceID != s.serviceID { continue } - if s.mu.lockAdmissions != 0 { - s.mu.pendingRefReleases = append(s.mu.pendingRefReleases, bind) - continue - } s.releaseBindRefLocked(bind.Group, bind.Table, bind, s.mu.drainSnapshotReady) } } @@ -465,21 +462,6 @@ func (s *service) releaseBindRefLocked( } } -func (s *service) applyPendingRefReleasesLocked() { - if s.mu.lockAdmissions != 0 || len(s.mu.pendingRefReleases) == 0 { - return - } - for _, bind := range s.mu.pendingRefReleases { - s.releaseBindRefLocked( - bind.Group, - bind.Table, - bind, - s.mu.drainSnapshotReady, - ) - } - s.mu.pendingRefReleases = s.mu.pendingRefReleases[:0] -} - func (s *service) checkCanMoveGroupTables() { s.mu.Lock() defer s.mu.Unlock() @@ -556,36 +538,74 @@ func (s *service) canLockOnServiceStatusLocked( return false } +type lockAdmission struct { + preDrain bool + reservedBind pb.LockTable + reserved bool +} + +func (a *lockAdmission) consume(bind pb.LockTable) bool { + if !a.reserved || + a.reservedBind.Group != bind.Group || + a.reservedBind.Table != bind.Table { + return false + } + a.reserved = false + return true +} + func (s *service) beginLockAdmission( txnID []byte, opts pb.LockOptions, tableID uint64, rows [][]byte, -) (bool, bool) { +) (lockAdmission, bool) { s.mu.Lock() defer s.mu.Unlock() if !s.canLockOnServiceStatusLocked(txnID, opts, tableID, rows) { - return false, false + return lockAdmission{}, false + } + admission := lockAdmission{ + preDrain: s.mu.status == pb.Status_ServiceLockEnable, + } + if !admission.preDrain { + if opts.Sharding == pb.Sharding_ByRow { + tableID = ShardingByRow(rows[0]) + } + l := s.tableGroups.get(opts.Group, tableID) + if l == nil { + return lockAdmission{}, false + } + admission.reservedBind = l.getBind() + admission.reserved = true + s.mu.lockTableRef[opts.Group][tableID]++ } - preDrain := s.mu.status == pb.Status_ServiceLockEnable s.mu.lockAdmissions++ - return true, preDrain + return admission, true } -func (s *service) endLockAdmission(preDrain bool) { +func (s *service) endLockAdmission(admission lockAdmission) { s.mu.Lock() defer s.mu.Unlock() if s.mu.lockAdmissions == 0 { panic("lock admission underflow") } s.mu.lockAdmissions-- - if preDrain && s.mu.status == pb.Status_ServiceLockWaiting { + if admission.preDrain && s.mu.status == pb.Status_ServiceLockWaiting { if s.mu.preDrainAdmissions == 0 { panic("pre-drain lock admission underflow") } s.mu.preDrainAdmissions-- } - s.applyPendingRefReleasesLocked() + if admission.reserved { + bind := admission.reservedBind + s.releaseBindRefLocked( + bind.Group, + bind.Table, + bind, + s.mu.drainSnapshotReady, + ) + } s.prepareDrainSnapshotLocked() s.tryCompleteDrainLocked() } @@ -593,7 +613,6 @@ func (s *service) endLockAdmission(preDrain bool) { func (s *service) tryCompleteDrain() { s.mu.Lock() defer s.mu.Unlock() - s.applyPendingRefReleasesLocked() s.prepareDrainSnapshotLocked() s.tryCompleteDrainLocked() } diff --git a/pkg/lockservice/service_remote.go b/pkg/lockservice/service_remote.go index e26bf459fc0c5..d8ee1af06b7e0 100644 --- a/pkg/lockservice/service_remote.go +++ b/pkg/lockservice/service_remote.go @@ -275,12 +275,12 @@ func (s *service) handleRemoteLock( resp *pb.Response, cs morpc.ClientSession) { logFields := remoteLockResponseLogFields(req) - admitted, preDrain := s.beginLockAdmission(req.Lock.TxnID, req.Lock.Options, req.LockTable.Table, req.Lock.Rows) + admission, admitted := s.beginLockAdmission(req.Lock.TxnID, req.Lock.Options, req.LockTable.Table, req.Lock.Rows) if !admitted { _ = writeResponseWithDeadline(s.logger, cancel, resp, moerr.NewRetryForCNRollingRestart(), cs, defaultRPCWriteTimeout, logFields) return } - defer s.endLockAdmission(preDrain) + defer func() { s.endLockAdmission(admission) }() l, err := s.getLocalLockTable(ctx, req, resp) if err != nil || @@ -334,7 +334,9 @@ func (s *service) handleRemoteLock( return } - if txn.lockTableBindTouched(bind) && bind.ServiceID == s.serviceID { + if txn.lockTableBindTouched(bind) && + bind.ServiceID == s.serviceID && + !admission.consume(bind) { s.incRef(bind.Group, bind.Table) } txnID := append([]byte(nil), req.Lock.TxnID...) @@ -369,12 +371,12 @@ func (s *service) handleForwardLock( resp *pb.Response, cs morpc.ClientSession) { logFields := remoteLockResponseLogFields(req) - admitted, preDrain := s.beginLockAdmission(req.Lock.TxnID, req.Lock.Options, req.LockTable.Table, req.Lock.Rows) + admission, admitted := s.beginLockAdmission(req.Lock.TxnID, req.Lock.Options, req.LockTable.Table, req.Lock.Rows) if !admitted { _ = writeResponseWithDeadline(s.logger, cancel, resp, moerr.NewRetryForCNRollingRestart(), cs, defaultRPCWriteTimeout, logFields) return } - defer s.endLockAdmission(preDrain) + defer func() { s.endLockAdmission(admission) }() l, err := s.getLockTable( ctx, @@ -426,7 +428,9 @@ func (s *service) handleForwardLock( return } - if txn.lockTableBindTouched(bind) && bind.ServiceID == s.serviceID { + if txn.lockTableBindTouched(bind) && + bind.ServiceID == s.serviceID && + !admission.consume(bind) { s.incRef(bind.Group, bind.Table) } txnID := append([]byte(nil), req.Lock.TxnID...) diff --git a/pkg/lockservice/service_test.go b/pkg/lockservice/service_test.go index e73ed0ed3a2ed..e6a995efc05e6 100644 --- a/pkg/lockservice/service_test.go +++ b/pkg/lockservice/service_test.go @@ -6052,15 +6052,15 @@ func TestPostGateAdmissionsDoNotBlockDrainSnapshot(t *testing.T) { _, err := s.Lock(context.Background(), table, newTestRows(1), txnID, options) require.NoError(t, err) - admitted, preDrain := s.beginLockAdmission(txnID, options, table, newTestRows(1)) + preDrain, admitted := s.beginLockAdmission(txnID, options, table, newTestRows(1)) require.True(t, admitted) - require.True(t, preDrain) + require.True(t, preDrain.preDrain) s.checkCanMoveGroupTables() for range 100 { - admitted, postDrain := s.beginLockAdmission(txnID, options, table, newTestRows(1)) + postDrain, admitted := s.beginLockAdmission(txnID, options, table, newTestRows(1)) require.True(t, admitted) - require.False(t, postDrain) + require.False(t, postDrain.preDrain) s.endLockAdmission(postDrain) } s.mu.RLock() @@ -6077,6 +6077,46 @@ func TestPostGateAdmissionsDoNotBlockDrainSnapshot(t *testing.T) { }) } +func TestLongLockWaitDoesNotRetainCompletedTxnRefs(t *testing.T) { + runLockServiceTests(t, []string{"s1"}, + func(_ *lockTableAllocator, services []*service) { + s := services[0] + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + options := newTestRowExclusiveOptions() + waitTable := uint64(257910) + row := []byte{1} + holderTxn := []byte("long-wait-holder") + waiterTxn := []byte("long-wait-waiter") + + _, err := s.Lock(ctx, waitTable, [][]byte{row}, holderTxn, options) + require.NoError(t, err) + waitDone := make(chan error, 1) + go func() { + _, err := s.Lock(ctx, waitTable, [][]byte{row}, waiterTxn, options) + waitDone <- err + }() + waitWaiters(t, s, waitTable, row, 1) + + for i := range 100 { + table := uint64(258000 + i) + txnID := []byte(fmt.Sprintf("completed-during-wait-%d", i)) + _, err := s.Lock(ctx, table, newTestRows(1), txnID, options) + require.NoError(t, err) + require.NoError(t, s.Unlock(ctx, txnID, timestamp.Timestamp{})) + require.False(t, s.validGroupTable(0, table), + "completed transaction ref must be released immediately") + } + s.mu.RLock() + require.Len(t, s.mu.lockTableRef[0], 1) + s.mu.RUnlock() + + require.NoError(t, s.Unlock(ctx, holderTxn, timestamp.Timestamp{})) + require.NoError(t, <-waitDone) + require.NoError(t, s.Unlock(ctx, waiterTxn, timestamp.Timestamp{})) + }) +} + func TestDrainWaitsForTxnCloseLinearization(t *testing.T) { runLockServiceTests(t, []string{"s1"}, func(_ *lockTableAllocator, services []*service) {