diff --git a/pkg/clusterservice/cluster.go b/pkg/clusterservice/cluster.go index c679ebd4f54ef..2c933656e81fa 100644 --- a/pkg/clusterservice/cluster.go +++ b/pkg/clusterservice/cluster.go @@ -121,6 +121,33 @@ func GetCNServiceWithoutWorkingStateWithContext( return ctx.Err() } +// GetAllTNServicesWithContext returns a TN service snapshot without waiting +// past ctx for the built-in cluster's initial HAKeeper refresh. +func GetAllTNServicesWithContext( + ctx context.Context, + service MOCluster, +) ([]metadata.TNService, error) { + if ctx == nil { + ctx = context.Background() + } + if err := ctx.Err(); err != nil { + return nil, err + } + if service == nil { + return nil, moerr.NewInternalErrorNoCtx("mocluster service is not initialized") + } + if builtIn, ok := service.(*cluster); ok { + if err := builtIn.waitReadyWithContext(ctx); err != nil { + return nil, err + } + services := builtIn.services.Load() + return append([]metadata.TNService(nil), services.tn...), ctx.Err() + } + + services := service.GetAllTNServices() + return services, ctx.Err() +} + func lookupMOCluster(service string) (MOCluster, bool, error) { rt := runtime.ServiceRuntime(service) if rt == nil { @@ -339,9 +366,14 @@ func (c *cluster) Refresh(ctx context.Context) error { } func (c *cluster) Close() { - c.waitReady() c.stopper.Stop() - close(c.forceRefreshC) + // A failed initial refresh leaves readiness waiters blocked. Once the + // refresh task has stopped, release them so shutdown does not depend on + // HAKeeper becoming available. + c.readyOnce.Do(func() { + c.ready.Store(true) + close(c.readyC) + }) } // DebugUpdateCNLabel implements the MOCluster interface. diff --git a/pkg/clusterservice/cluster_test.go b/pkg/clusterservice/cluster_test.go index f6ca7f052119a..d32b29c93ee74 100644 --- a/pkg/clusterservice/cluster_test.go +++ b/pkg/clusterservice/cluster_test.go @@ -70,6 +70,16 @@ func TestCNServiceSnapshotHonorsCancellationWhileClusterStarts(t *testing.T) { require.ErrorIs(t, err, context.DeadlineExceeded) } +func TestTNServiceSnapshotHonorsCancellationWhileClusterStarts(t *testing.T) { + c := &cluster{readyC: make(chan struct{})} + c.services.Store(&services{}) + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond) + defer cancel() + + _, err := GetAllTNServicesWithContext(ctx, c) + require.ErrorIs(t, err, context.DeadlineExceeded) +} + func TestClusterForceRefresh(t *testing.T) { runClusterTest( time.Hour, diff --git a/pkg/tnservice/replica.go b/pkg/tnservice/replica.go index 21b0f5264ac00..c9542065e69f6 100644 --- a/pkg/tnservice/replica.go +++ b/pkg/tnservice/replica.go @@ -35,6 +35,7 @@ type replica struct { logger *log.MOLogger shard metadata.TNShard service service.TxnService + serviceC chan struct{} startedC chan struct{} createCtx context.Context cancelCreate context.CancelFunc @@ -81,6 +82,7 @@ func newReplica(shard metadata.TNShard, rt runtime.Runtime) *replica { rt: rt, shard: shard, logger: rt.Logger().With(util.TxnTNShardField(shard)), + serviceC: make(chan struct{}), startedC: make(chan struct{}), createCtx: ctx, cancelCreate: cancel, @@ -114,6 +116,7 @@ func (r *replica) startReserved(txnService service.TxnService) error { } r.service = txnService r.mu.Unlock() + close(r.serviceC) err := txnService.Start() r.finishStart(err) @@ -159,6 +162,30 @@ func (r *replica) close(destroy bool) error { return r.closeErr } +func (r *replica) cancelRecovery() { + r.mu.RLock() + starting := r.mu.starting + txnService := r.service + r.mu.RUnlock() + if !starting { + return + } + if txnService == nil { + // Once start is reserved, startReserved either publishes the service or + // finishStart reports that startup ended without one. + select { + case <-r.serviceC: + case <-r.startedC: + } + r.mu.RLock() + txnService = r.service + r.mu.RUnlock() + } + if txnService != nil { + txnService.CancelRecovery() + } +} + func (r *replica) closeOnceFn() error { r.mu.RLock() starting := r.mu.starting @@ -166,6 +193,8 @@ func (r *replica) closeOnceFn() error { if !starting { return nil } + // Recovery may block Start indefinitely while waiting for a participant. + r.cancelRecovery() r.waitStartCompleted() r.mu.Lock() diff --git a/pkg/tnservice/replica_test.go b/pkg/tnservice/replica_test.go index 7b4b471f2c23c..bb1dc6c3c9836 100644 --- a/pkg/tnservice/replica_test.go +++ b/pkg/tnservice/replica_test.go @@ -17,13 +17,20 @@ package tnservice import ( "context" "errors" + "sync" "testing" "time" + "github.com/matrixorigin/matrixone/pkg/clusterservice" "github.com/matrixorigin/matrixone/pkg/common/runtime" + "github.com/matrixorigin/matrixone/pkg/defines" + "github.com/matrixorigin/matrixone/pkg/fileservice" + logpb "github.com/matrixorigin/matrixone/pkg/pb/logservice" + "github.com/matrixorigin/matrixone/pkg/pb/metadata" "github.com/matrixorigin/matrixone/pkg/pb/txn" "github.com/matrixorigin/matrixone/pkg/txn/service" "github.com/matrixorigin/matrixone/pkg/txn/storage" + "github.com/matrixorigin/matrixone/pkg/txn/storage/mem" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -39,6 +46,39 @@ type startErrorStorage struct { destroyCalls int } +type closeUnblocksStartTxnService struct { + service.TxnService + started chan struct{} + closed chan struct{} + startOnce sync.Once + closeOnce sync.Once +} + +type signalingRecoveryCluster struct { + clusterservice.MOCluster + entered chan struct{} + once sync.Once +} + +func (c *signalingRecoveryCluster) GetAllTNServices() []metadata.TNService { + c.once.Do(func() { close(c.entered) }) + return c.MOCluster.GetAllTNServices() +} + +func (s *closeUnblocksStartTxnService) Start() error { + s.startOnce.Do(func() { close(s.started) }) + <-s.closed + return context.Canceled +} + +func (s *closeUnblocksStartTxnService) CancelRecovery() { + s.closeOnce.Do(func() { close(s.closed) }) +} + +func (s *closeUnblocksStartTxnService) Close(bool) error { + return nil +} + type closeTrackingTxnService struct { service.TxnService closeCalls int @@ -210,6 +250,133 @@ func TestCloseFailedStartReplica(t *testing.T) { } } +func TestCloseCancelsBlockedReplicaStart(t *testing.T) { + txnService := &closeUnblocksStartTxnService{ + started: make(chan struct{}), + closed: make(chan struct{}), + } + r := newReplica(newTestTNShard(1, 2, 3), runtime.DefaultRuntime()) + startResult := make(chan error, 1) + go func() { + startResult <- r.start(txnService) + }() + + select { + case <-txnService.started: + case <-time.After(time.Second): + t.Fatal("replica start did not begin") + } + + closeResult := make(chan error, 1) + go func() { + closeResult <- r.close(false) + }() + + select { + case err := <-closeResult: + require.ErrorIs(t, err, context.Canceled) + case <-time.After(time.Second): + t.Fatal("close did not cancel blocked replica start") + } + require.ErrorIs(t, <-startResult, context.Canceled) +} + +func TestRemoveReplicaCancelsBlockedStartBeforeWaiting(t *testing.T) { + txnService := &closeUnblocksStartTxnService{ + started: make(chan struct{}), + closed: make(chan struct{}), + } + r := newReplica(newTestTNShard(1, 2, 3), runtime.DefaultRuntime()) + startResult := make(chan error, 1) + go func() { startResult <- r.start(txnService) }() + select { + case <-txnService.started: + case <-time.After(time.Second): + t.Fatal("replica start did not begin") + } + + fs, err := fileservice.NewMemoryFS( + defines.LocalFileServiceName, + fileservice.DisabledCacheConfig, + nil, + ) + require.NoError(t, err) + s := &store{ + cfg: &Config{UUID: "test"}, + rt: runtime.DefaultRuntime(), + metadataFileService: fs, + replicas: &sync.Map{}, + } + s.replicas.Store(r.shard.ShardID, r) + + removed := make(chan error, 1) + go func() { removed <- s.removeReplicaLocked(r.shard.ShardID) }() + select { + case err := <-removed: + require.ErrorIs(t, err, context.Canceled) + case <-time.After(time.Second): + t.Fatal("removeReplicaLocked waited for Start before canceling recovery") + } + require.ErrorIs(t, <-startResult, context.Canceled) + require.Nil(t, s.getReplica(r.shard.ShardID)) +} + +func TestCloseCancelsReplicaBlockedInRecovery(t *testing.T) { + meta := service.NewTestTxn(1, 1, 1) + meta.Status = txn.TxnStatus_Prepared + meta.PreparedTS = service.NewTestTimestamp(2) + meta.TNShards = append(meta.TNShards, metadata.TNShard{ + TNShardRecord: metadata.TNShardRecord{ShardID: 99}, + }) + mlog := mem.NewMemLog() + data := (&mem.KVLog{Txn: meta}).MustMarshal() + record := mlog.GetLogRecord(len(data)) + record.Type = logpb.UserRecord + record.Data = data + _, err := mlog.Append(context.Background(), record) + require.NoError(t, err) + + sender := service.NewTestSender() + t.Cleanup(func() { require.NoError(t, sender.Close()) }) + txnService := service.NewTestTxnServiceWithLog( + t, 1, sender, service.NewTestClock(0), mlog) + baseCluster := clusterservice.NewMOCluster( + "dn-uuid", nil, time.Hour, + clusterservice.WithDisableRefresh(), + clusterservice.WithServices(nil, nil), + ) + t.Cleanup(baseCluster.Close) + cluster := &signalingRecoveryCluster{ + MOCluster: baseCluster, + entered: make(chan struct{}), + } + runtime.ServiceRuntime("dn-uuid").SetGlobalVariables(runtime.ClusterService, cluster) + + r := newReplica(newTestTNShard(1, 2, 3), runtime.DefaultRuntime()) + startResult := make(chan error, 1) + go func() { startResult <- r.start(txnService) }() + select { + case <-cluster.entered: + case <-time.After(time.Second): + t.Fatal("recovery did not reach the missing participant route wait") + } + + closeResult := make(chan error, 1) + go func() { closeResult <- r.close(false) }() + select { + case err := <-closeResult: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("replica close did not cancel real transaction recovery") + } + select { + case err := <-startResult: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("replica start remained blocked after recovery cancellation") + } +} + func TestWaitStarted(t *testing.T) { r := newReplica(newTestTNShard(1, 2, 3), runtime.DefaultRuntime()) c := make(chan struct{}) diff --git a/pkg/tnservice/store.go b/pkg/tnservice/store.go index f376e37a514af..c13b28f7c53d8 100644 --- a/pkg/tnservice/store.go +++ b/pkg/tnservice/store.go @@ -238,16 +238,20 @@ func (s *store) Start() error { } func (s *store) Close() error { - s.stopper.Stop() - s.moCluster.Close() - - var err error // Reject new replica calls and cancel active call contexts before waiting - // for the RPC server to drain. Storage remains open until the drain ends. + // for store tasks. A published service may be blocked in recovery, so its + // cancellation must be delivered before joining the store stopper. Storage + // remains open until the RPC server drains below. s.replicas.Range(func(_, value any) bool { - value.(*replica).cancelStart(false) + r := value.(*replica) + r.cancelStart(false) + r.cancelRecovery() return true }) + s.stopper.Stop() + s.moCluster.Close() + + var err error if s.queryService != nil { err = errors.Join(err, s.queryService.Close()) } @@ -333,9 +337,7 @@ func (s *store) createReplicaLocked(shard metadata.TNShard) error { } err := s.stopper.RunTask(func(stopperCtx context.Context) { - stopCancelPropagation := context.AfterFunc(stopperCtx, func() { - r.cancelStart(false) - }) + stopCancelPropagation := propagateReplicaStopperCancellation(stopperCtx, r) defer stopCancelPropagation() for { @@ -403,6 +405,16 @@ func (s *store) createReplicaLocked(shard metadata.TNShard) error { return nil } +func propagateReplicaStopperCancellation( + stopperCtx context.Context, + r *replica, +) func() bool { + return context.AfterFunc(stopperCtx, func() { + r.cancelStart(false) + r.cancelRecovery() + }) +} + func waitCreateRetry(stopperCtx, createCtx context.Context) error { timer := time.NewTimer(retryCreateStorageInterval) defer timer.Stop() @@ -418,8 +430,6 @@ func waitCreateRetry(stopperCtx, createCtx context.Context) error { func (s *store) removeReplicaLocked(tnShardID uint64) error { if r := s.getReplica(tnShardID); r != nil { - r.cancelStart(true) - r.waitStartCompleted() err := r.close(true) s.replicas.CompareAndDelete(tnShardID, r) s.removeTNShardLocked(tnShardID) diff --git a/pkg/tnservice/store_rpc_handler_test.go b/pkg/tnservice/store_rpc_handler_test.go index 1b1d67256d3b8..fc26c52409c22 100644 --- a/pkg/tnservice/store_rpc_handler_test.go +++ b/pkg/tnservice/store_rpc_handler_test.go @@ -47,6 +47,8 @@ func (s *leaseCancelReadTxnService) Start() error { return nil } +func (s *leaseCancelReadTxnService) CancelRecovery() {} + func (s *leaseCancelReadTxnService) Close(bool) error { return nil } diff --git a/pkg/tnservice/store_test.go b/pkg/tnservice/store_test.go index 0b8e3110dfc7d..473139c97bdc8 100644 --- a/pkg/tnservice/store_test.go +++ b/pkg/tnservice/store_test.go @@ -24,6 +24,7 @@ import ( "testing" "time" + "github.com/matrixorigin/matrixone/pkg/clusterservice" "github.com/matrixorigin/matrixone/pkg/common/runtime" "github.com/matrixorigin/matrixone/pkg/defines" "github.com/matrixorigin/matrixone/pkg/fileservice" @@ -31,6 +32,7 @@ import ( "github.com/matrixorigin/matrixone/pkg/logutil" logservicepb "github.com/matrixorigin/matrixone/pkg/pb/logservice" "github.com/matrixorigin/matrixone/pkg/pb/metadata" + "github.com/matrixorigin/matrixone/pkg/pb/txn" "github.com/matrixorigin/matrixone/pkg/queryservice" "github.com/matrixorigin/matrixone/pkg/queryservice/client" "github.com/matrixorigin/matrixone/pkg/txn/clock" @@ -307,6 +309,238 @@ func TestStoreCloseCancelsReplicasBeforeDrainingRPCServer(t *testing.T) { require.True(t, canceledBeforeServerClose.Load()) } +func TestStoreCloseCancelsRecoveryBeforeStoppingReplicaTask(t *testing.T) { + runtime.SetupServiceBasedRuntime("", runtime.DefaultRuntime()) + runtime.SetupServiceBasedRuntime("u1", runtime.ServiceRuntime("")) + fsFactory := func(name string) (*fileservice.FileServices, error) { + fs, err := fileservice.NewMemoryFS(name, fileservice.DisabledCacheConfig, nil) + if err != nil { + return nil, err + } + return fileservice.NewFileServices(name, fs) + } + s := newTestStore( + t, + "u1", + fsFactory, + WithHAKeeperClientFactory(func() (logservice.TNHAKeeperClient, error) { + return newTestHAKeeperClient(), nil + }), + WithLogServiceClientFactory(func(metadata.TNShard) (logservice.Client, error) { + return mem.NewMemLog(), nil + }), + WithConfigAdjust(func(c *Config) { + c.HAKeeper.HeatbeatInterval.Duration = 10 * time.Millisecond + c.Txn.Storage.Backend = StorageMEMKV + }), + ) + require.NoError(t, s.Start()) + + meta := service.NewTestTxn(1, 1, 1) + meta.Status = txn.TxnStatus_Prepared + meta.PreparedTS = service.NewTestTimestamp(2) + meta.TNShards = append(meta.TNShards, metadata.TNShard{ + TNShardRecord: metadata.TNShardRecord{ShardID: 99}, + }) + mlog := mem.NewMemLog() + data := (&mem.KVLog{Txn: meta}).MustMarshal() + record := mlog.GetLogRecord(len(data)) + record.Type = logservicepb.UserRecord + record.Data = data + _, err := mlog.Append(context.Background(), record) + require.NoError(t, err) + + sender := service.NewTestSender() + t.Cleanup(func() { require.NoError(t, sender.Close()) }) + txnService := service.NewTestTxnServiceWithLog( + t, 1, sender, service.NewTestClock(0), mlog) + baseCluster := clusterservice.NewMOCluster( + "dn-uuid", nil, time.Hour, + clusterservice.WithDisableRefresh(), + clusterservice.WithServices(nil, nil), + ) + t.Cleanup(baseCluster.Close) + cluster := &signalingRecoveryCluster{ + MOCluster: baseCluster, + entered: make(chan struct{}), + } + runtime.ServiceRuntime("dn-uuid").SetGlobalVariables(runtime.ClusterService, cluster) + + r := newReplica(newTestTNShard(1, 2, 3), s.rt) + s.replicas.Store(r.shard.ShardID, r) + require.NoError(t, s.stopper.RunTask(func(context.Context) { + _ = r.start(txnService) + })) + select { + case <-cluster.entered: + case <-time.After(time.Second): + t.Fatal("recovery did not reach the missing participant route wait") + } + + closed := make(chan error, 1) + go func() { closed <- s.Close() }() + select { + case err := <-closed: + require.NoError(t, err) + case <-time.After(time.Second): + require.NoError(t, r.close(false)) + require.NoError(t, <-closed) + t.Fatal("store Close waited for the stopper before canceling recovery") + } +} + +type neverReadyHAKeeperClient struct { + *testHAKeeperClient + called chan struct{} + once sync.Once +} + +func (c *neverReadyHAKeeperClient) GetClusterDetails( + context.Context, +) (logservicepb.ClusterDetails, error) { + c.once.Do(func() { close(c.called) }) + return logservicepb.ClusterDetails{}, errors.New("injected hakeeper failure") +} + +func TestStoreCloseWhenClusterNeverBecomesReady(t *testing.T) { + runtime.SetupServiceBasedRuntime("", runtime.DefaultRuntime()) + runtime.SetupServiceBasedRuntime("u1", runtime.ServiceRuntime("")) + fsFactory := func(name string) (*fileservice.FileServices, error) { + fs, err := fileservice.NewMemoryFS(name, fileservice.DisabledCacheConfig, nil) + if err != nil { + return nil, err + } + return fileservice.NewFileServices(name, fs) + } + hakeeper := &neverReadyHAKeeperClient{ + testHAKeeperClient: newTestHAKeeperClient(), + called: make(chan struct{}), + } + s := newTestStore( + t, + "u1", + fsFactory, + WithHAKeeperClientFactory(func() (logservice.TNHAKeeperClient, error) { + return hakeeper, nil + }), + WithLogServiceClientFactory(func(metadata.TNShard) (logservice.Client, error) { + return mem.NewMemLog(), nil + }), + WithConfigAdjust(func(c *Config) { + c.Cluster.RefreshInterval.Duration = 10 * time.Millisecond + c.HAKeeper.HeatbeatInterval.Duration = 10 * time.Millisecond + c.Txn.Storage.Backend = StorageMEMKV + }), + ) + require.NoError(t, s.Start()) + select { + case <-hakeeper.called: + case <-time.After(time.Second): + t.Fatal("cluster refresh did not start") + } + + closed := make(chan error, 1) + go func() { closed <- s.Close() }() + select { + case err := <-closed: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("store Close waited forever for cluster readiness") + } +} + +type blockingCancelRecoveryTxnService struct { + service.TxnService + entered chan struct{} + release chan struct{} + once sync.Once +} + +func (s *blockingCancelRecoveryTxnService) Start() error { + return nil +} + +func (s *blockingCancelRecoveryTxnService) CancelRecovery() { + s.once.Do(func() { close(s.entered) }) + <-s.release +} + +func (s *blockingCancelRecoveryTxnService) Close(bool) error { + return nil +} + +func TestStoreCloseCancelsLateReplicaRecoveryFromStopper(t *testing.T) { + runtime.SetupServiceBasedRuntime("", runtime.DefaultRuntime()) + runtime.SetupServiceBasedRuntime("u1", runtime.ServiceRuntime("")) + fsFactory := func(name string) (*fileservice.FileServices, error) { + fs, err := fileservice.NewMemoryFS(name, fileservice.DisabledCacheConfig, nil) + if err != nil { + return nil, err + } + return fileservice.NewFileServices(name, fs) + } + s := newTestStore( + t, + "u1", + fsFactory, + WithHAKeeperClientFactory(func() (logservice.TNHAKeeperClient, error) { + return newTestHAKeeperClient(), nil + }), + WithLogServiceClientFactory(func(metadata.TNShard) (logservice.Client, error) { + return mem.NewMemLog(), nil + }), + WithConfigAdjust(func(c *Config) { + c.HAKeeper.HeatbeatInterval.Duration = 10 * time.Millisecond + c.Txn.Storage.Backend = StorageMEMKV + }), + ) + require.NoError(t, s.Start()) + + shard := newTestTNShard(1, 2, 3) + barrierService := &blockingCancelRecoveryTxnService{ + entered: make(chan struct{}), + release: make(chan struct{}), + } + initial := newReplica(shard, s.rt) + require.NoError(t, initial.start(barrierService)) + s.replicas.Store(shard.ShardID, initial) + + closed := make(chan error, 1) + go func() { closed <- s.Close() }() + select { + case <-barrierService.entered: + case <-time.After(time.Second): + t.Fatal("store Close did not enter the initial replica cancellation") + } + + lateService := &closeUnblocksStartTxnService{ + started: make(chan struct{}), + closed: make(chan struct{}), + } + late := newReplica(shard, s.rt) + require.True(t, s.replicas.CompareAndSwap(shard.ShardID, initial, late)) + require.NoError(t, s.stopper.RunTask(func(stopperCtx context.Context) { + stopCancelPropagation := propagateReplicaStopperCancellation(stopperCtx, late) + defer stopCancelPropagation() + _ = late.start(lateService) + })) + select { + case <-lateService.started: + case <-time.After(time.Second): + t.Fatal("late replica did not block in Start") + } + + close(barrierService.release) + select { + case err := <-closed: + require.ErrorIs(t, err, context.Canceled) + case <-time.After(time.Second): + late.cancelRecovery() + require.ErrorIs(t, <-closed, context.Canceled) + t.Fatal("store Close did not cancel recovery for a replica registered after its initial Range") + } +} + func TestHeartbeatOnlyReportsStartedReplicas(t *testing.T) { runTNStoreTest(t, func(s *store) { r := newReplica(newTestTNShard(1, 2, 3), s.rt) diff --git a/pkg/txn/service/service.go b/pkg/txn/service/service.go index 8565ea203d3f5..460b25b3906d9 100644 --- a/pkg/txn/service/service.go +++ b/pkg/txn/service/service.go @@ -60,12 +60,14 @@ type service struct { // due to the network, resulting in the transaction information being written back to the map. // We use the zombieTimeout setting to solve this problem, so that when a transaction exceeds the zombieTimeout // threshold in the map, it is cleaned up. - transactions sync.Map // string(txn.id) -> txnContext - zombieTimeout time.Duration - pool sync.Pool - recoveryC chan struct{} - recoveryOnce sync.Once - txnC chan txn.TxnMeta + transactions sync.Map // string(txn.id) -> txnContext + zombieTimeout time.Duration + pool sync.Pool + recoveryC chan struct{} + recoveryOnce sync.Once + recoveryCtx context.Context + recoveryCancel context.CancelFunc + txnC chan txn.TxnMeta } // NewTxnService create TxnService @@ -78,6 +80,7 @@ func NewTxnService( allocator lockservice.LockTableAllocator, ) TxnService { logger := util.GetLogger(sid) + recoveryCtx, recoveryCancel := context.WithCancel(context.Background()) s := &service{ sid: sid, logger: logger, @@ -92,10 +95,12 @@ func NewTxnService( shard.ShardID, shard.ReplicaID), stopper.WithLogger(logger.RawLogger())), - zombieTimeout: zombieTimeout, - recoveryC: make(chan struct{}), - txnC: make(chan txn.TxnMeta, 16), - allocator: allocator, + zombieTimeout: zombieTimeout, + recoveryC: make(chan struct{}), + recoveryCtx: recoveryCtx, + recoveryCancel: recoveryCancel, + txnC: make(chan txn.TxnMeta, 16), + allocator: allocator, } if err := s.stopper.RunTask(s.gcZombieTxn); err != nil { s.logger.Fatal("start gc zombie txn failed", @@ -119,7 +124,15 @@ func (s *service) Start() error { return nil } +// CancelRecovery interrupts recovery without closing storage. TN replica +// shutdown uses this to unblock Start before running the normal Close path. +func (s *service) CancelRecovery() { + s.recoveryCancel() +} + func (s *service) Close(destroy bool) error { + s.CancelRecovery() + s.finishRecovery() s.waitRecoveryCompleted() s.stopper.Stop() closer := s.storage.Close diff --git a/pkg/txn/service/service_recovery.go b/pkg/txn/service/service_recovery.go index 4ddde85db5bce..2b34c3ed909fe 100644 --- a/pkg/txn/service/service_recovery.go +++ b/pkg/txn/service/service_recovery.go @@ -16,7 +16,10 @@ package service import ( "context" + "time" + "github.com/matrixorigin/matrixone/pkg/clusterservice" + "github.com/matrixorigin/matrixone/pkg/pb/metadata" "github.com/matrixorigin/matrixone/pkg/pb/txn" "github.com/matrixorigin/matrixone/pkg/txn/util" "go.uber.org/zap" @@ -27,15 +30,18 @@ func (s *service) startRecovery() { s.logger.Fatal("start recover task failed", zap.Error(err)) } - s.storage.StartRecovery(context.TODO(), s.txnC) + s.storage.StartRecovery(s.recoveryCtx, s.txnC) s.waitRecoveryCompleted() } func (s *service) doRecovery(ctx context.Context) { + defer s.finishRecovery() for { select { case <-ctx.Done(): return + case <-s.recoveryCtx.Done(): + return case txn, ok := <-s.txnC: if !ok { s.end() @@ -46,6 +52,130 @@ func (s *service) doRecovery(ctx context.Context) { } } +type recoveryRouteResolver struct { + cluster clusterservice.MOCluster + services []metadata.TNService + initialized bool +} + +func (r *recoveryRouteResolver) refresh( + ctx context.Context, + s *service, +) bool { + const retryInterval = 100 * time.Millisecond + for r.cluster == nil { + var err error + r.cluster, err = clusterservice.GetMOClusterWithContext(ctx, s.sid) + if err == nil { + break + } + if !waitRecoveryRouteRetry(ctx, retryInterval) { + return false + } + } + + for { + // A complete cached route can still name a replaced replica. Establish + // freshness before accepting ReplicaID and Address for recovery RPCs. + if refresher, ok := r.cluster.(clusterservice.AuthoritativeRefresher); ok { + if err := refresher.Refresh(ctx); err != nil { + if ctx.Err() != nil || !waitRecoveryRouteRetry(ctx, retryInterval) { + return false + } + continue + } + } + services, err := clusterservice.GetAllTNServicesWithContext(ctx, r.cluster) + if err != nil { + if !waitRecoveryRouteRetry(ctx, retryInterval) { + return false + } + continue + } + r.services = services + r.initialized = true + return true + } +} + +func (s *service) resolveRecoveryTNShards( + ctx context.Context, + txnMeta txn.TxnMeta, + resolver *recoveryRouteResolver, +) (txn.TxnMeta, bool) { + const retryInterval = 100 * time.Millisecond + if !resolver.initialized && !resolver.refresh(ctx, s) { + return txnMeta, false + } + + for { + if shards, ok := s.resolveRecoveryTNShardsFromSnapshot(txnMeta.TNShards, resolver.services); ok { + txnMeta.TNShards = shards + return txnMeta, true + } + + if !waitRecoveryRouteRetry(ctx, retryInterval) || !resolver.refresh(ctx, s) { + return txnMeta, false + } + } +} + +func waitRecoveryRouteRetry(ctx context.Context, interval time.Duration) bool { + timer := time.NewTimer(interval) + defer timer.Stop() + select { + case <-ctx.Done(): + return false + case <-timer.C: + return true + } +} + +func (s *service) resolveRecoveryTNShardsFromSnapshot( + participants []metadata.TNShard, + services []metadata.TNService, +) ([]metadata.TNShard, bool) { + routes := make(map[uint64]metadata.TNShard) + for _, service := range services { + for _, shard := range service.Shards { + shard.Address = service.TxnServiceAddress + if recoveryTNShardRouteComplete(shard) { + routes[shard.ShardID] = shard + } + } + } + + resolved := make([]metadata.TNShard, len(participants)) + for i, participant := range participants { + if participant.ShardID == s.shard.ShardID { + if !recoveryTNShardRouteComplete(s.shard) { + return nil, false + } + resolved[i] = s.shard + continue + } + route, ok := routes[participant.ShardID] + if !ok { + return nil, false + } + if route.LogShardID == 0 { + route.LogShardID = participant.LogShardID + } + resolved[i] = route + } + return resolved, true +} + +func recoveryTNShardRouteComplete(shard metadata.TNShard) bool { + // LogShardID is storage metadata, not part of txn RPC routing or TN shard + // identity. Dynamic HAKeeper TN snapshots currently expose only ShardID and + // ReplicaID, so preserve LogShardID when supplied but do not wait forever + // when a current route legitimately has it unset. + return shard.ShardID != 0 && + shard.ReplicaID != 0 && + shard.Address != "" +} + func (s *service) addLog(txnMeta txn.TxnMeta) { if len(txnMeta.TNShards) <= 1 { return @@ -91,12 +221,22 @@ func (s *service) addLog(txnMeta txn.TxnMeta) { func (s *service) end() { defer s.finishRecovery() + resolver := new(recoveryRouteResolver) s.transactions.Range(func(_, value any) bool { txnCtx := value.(*txnContext) txnMeta := txnCtx.getTxn() - if !s.shard.Equal(txnMeta.TNShards[0]) { + if len(txnMeta.TNShards) == 0 || + txnMeta.TNShards[0].ShardID != s.shard.ShardID { return true } + // Fold the complete log stream before waiting for live routes. A later + // terminal record may remove an obsolete prepared transaction entirely. + var resolved bool + txnMeta, resolved = s.resolveRecoveryTNShards(s.recoveryCtx, txnMeta, resolver) + if !resolved { + return false + } + txnCtx.updateTxn(txnMeta) switch txnMeta.Status { case txn.TxnStatus_Prepared: @@ -106,6 +246,7 @@ func (s *service) end() { util.TxnField(txnMeta)) } case txn.TxnStatus_Committing: + s.validTNShard(txnMeta.TNShards[0]) if err := s.startAsyncCommitTask(txnCtx); err != nil { s.logger.Error("start commit task failed during recovery", zap.Error(err), @@ -180,7 +321,4 @@ func (s *service) checkRecoveryStatus(txnMeta txn.TxnMeta) { util.TxnField(txnMeta)) } - if txnMeta.Status == txn.TxnStatus_Committing { - s.validTNShard(txnMeta.TNShards[0]) - } } diff --git a/pkg/txn/service/service_recovery_test.go b/pkg/txn/service/service_recovery_test.go index a0958b5dc192f..64ddcba504bfc 100644 --- a/pkg/txn/service/service_recovery_test.go +++ b/pkg/txn/service/service_recovery_test.go @@ -16,15 +16,521 @@ package service import ( "context" + "errors" + "sync" "testing" "time" + "github.com/matrixorigin/matrixone/pkg/clusterservice" + "github.com/matrixorigin/matrixone/pkg/common/runtime" "github.com/matrixorigin/matrixone/pkg/logservice" + logpb "github.com/matrixorigin/matrixone/pkg/pb/logservice" + "github.com/matrixorigin/matrixone/pkg/pb/metadata" "github.com/matrixorigin/matrixone/pkg/pb/txn" + "github.com/matrixorigin/matrixone/pkg/txn/rpc" "github.com/matrixorigin/matrixone/pkg/txn/storage/mem" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) +type recoveryRouteSender struct { + mu sync.Mutex + requests [][]txn.TxnRequest + notifyC chan struct{} + block bool +} + +func newRecoveryRouteSender() *recoveryRouteSender { + return &recoveryRouteSender{notifyC: make(chan struct{}, 1)} +} + +func (s *recoveryRouteSender) Send(ctx context.Context, requests []txn.TxnRequest) (*rpc.SendResult, error) { + copied := append([]txn.TxnRequest(nil), requests...) + s.mu.Lock() + s.requests = append(s.requests, copied) + s.mu.Unlock() + responses := make([]txn.TxnResponse, len(requests)) + for i := range responses { + meta := requests[i].Txn + meta.Status = txn.TxnStatus_Aborted + responses[i].Txn = &meta + } + select { + case s.notifyC <- struct{}{}: + default: + } + if s.block { + <-ctx.Done() + return nil, ctx.Err() + } + return &rpc.SendResult{Responses: responses}, nil +} + +func (s *recoveryRouteSender) Close() error { return nil } + +func installRecoveryCluster(t *testing.T, sid string, services ...metadata.TNService) { + t.Helper() + c := clusterservice.NewMOCluster( + sid, + nil, + time.Hour, + clusterservice.WithDisableRefresh(), + clusterservice.WithServices(nil, services), + ) + t.Cleanup(c.Close) + runtime.ServiceRuntime(sid).SetGlobalVariables(runtime.ClusterService, c) +} + +func installRecoveryRoutesForServices(t *testing.T, services ...*service) { + t.Helper() + require.NotEmpty(t, services) + tnServices := make([]metadata.TNService, 0, len(services)) + for _, service := range services { + tnServices = append(tnServices, metadata.TNService{ + TxnServiceAddress: service.shard.Address, + Shards: []metadata.TNShard{service.shard}, + }) + } + installRecoveryCluster(t, services[0].sid, tnServices...) +} + +type recoveryClusterClient struct { + mu sync.RWMutex + details logpb.ClusterDetails +} + +type staleRecoveryCluster struct { + clusterservice.MOCluster + mu sync.Mutex + stale []metadata.TNService + current []metadata.TNService + refreshed bool + refreshCalls int + refreshFails int +} + +func (c *staleRecoveryCluster) GetAllTNServices() []metadata.TNService { + c.mu.Lock() + defer c.mu.Unlock() + if c.refreshed { + return c.current + } + return c.stale +} + +func (c *staleRecoveryCluster) Refresh(ctx context.Context) error { + if err := ctx.Err(); err != nil { + return err + } + c.mu.Lock() + defer c.mu.Unlock() + c.refreshCalls++ + if c.refreshFails > 0 { + c.refreshFails-- + return errors.New("transient refresh failure") + } + c.refreshed = true + return nil +} + +func (c *recoveryClusterClient) GetClusterDetails(context.Context) (logpb.ClusterDetails, error) { + c.mu.RLock() + defer c.mu.RUnlock() + return c.details, nil +} + +func (c *recoveryClusterClient) setTNRoute(shardID, replicaID uint64, address string) { + c.mu.Lock() + defer c.mu.Unlock() + c.details.TNStores = []logpb.TNStore{{ + UUID: "tn", + ServiceAddress: address, + Shards: []logpb.TNShardInfo{{ + ShardID: shardID, + ReplicaID: replicaID, + }}, + }} +} + +func TestRecoveryRestoresRemoteParticipantRoutes(t *testing.T) { + sender := newRecoveryRouteSender() + s := NewTestTxnServiceWithLog(t, 1, sender, NewTestClock(0), nil).(*service) + defer s.stopper.Stop() + remote := metadata.TNShard{ + TNShardRecord: metadata.TNShardRecord{ShardID: 2}, + ReplicaID: 200, + } + installRecoveryCluster(t, s.sid, metadata.TNService{ + TxnServiceAddress: "tn-2", + Shards: []metadata.TNShard{remote}, + }) + + meta := NewTestTxn(1, 1, 1) + meta.Status = txn.TxnStatus_Prepared + meta.PreparedTS = NewTestTimestamp(2) + meta.TNShards = append(meta.TNShards, metadata.TNShard{ + TNShardRecord: metadata.TNShardRecord{ShardID: 2, LogShardID: 20}, + }) + + require.NoError(t, s.stopper.RunTask(s.doRecovery)) + s.txnC <- meta + close(s.txnC) + s.waitRecoveryCompleted() + + select { + case <-sender.notifyC: + case <-time.After(time.Second): + t.Fatal("recovery did not query the remote participant") + } + sender.mu.Lock() + defer sender.mu.Unlock() + require.NotEmpty(t, sender.requests) + require.Len(t, sender.requests[0], 1) + assert.Equal(t, metadata.TNShard{ + TNShardRecord: metadata.TNShardRecord{ShardID: 2, LogShardID: 20}, + ReplicaID: 200, + Address: "tn-2", + }, sender.requests[0][0].GetTargetTN()) +} + +func TestRecoveryRefreshesStaleParticipantRoutes(t *testing.T) { + for _, test := range []struct { + name string + persistedRoute metadata.TNShard + refreshFails int + expectRefreshes int + }{ + { + name: "incomplete persisted route with complete stale cache", + expectRefreshes: 1, + persistedRoute: metadata.TNShard{ + TNShardRecord: metadata.TNShardRecord{ShardID: 2}, + }, + }, + { + name: "complete stale persisted route", + expectRefreshes: 1, + persistedRoute: metadata.TNShard{ + TNShardRecord: metadata.TNShardRecord{ShardID: 2}, + ReplicaID: 100, + Address: "old-tn-2", + }, + }, + { + name: "transient authoritative refresh failure", + refreshFails: 1, + expectRefreshes: 2, + persistedRoute: metadata.TNShard{ + TNShardRecord: metadata.TNShardRecord{ShardID: 2}, + }, + }, + } { + t.Run(test.name, func(t *testing.T) { + sender := newRecoveryRouteSender() + s := NewTestTxnServiceWithLog(t, 1, sender, NewTestClock(0), nil).(*service) + defer s.stopper.Stop() + cluster := &staleRecoveryCluster{ + stale: []metadata.TNService{{ + TxnServiceAddress: "old-tn-2", + Shards: []metadata.TNShard{{ + TNShardRecord: metadata.TNShardRecord{ShardID: 2}, + ReplicaID: 100, + }}, + }}, + current: []metadata.TNService{{ + TxnServiceAddress: "current-tn-2", + Shards: []metadata.TNShard{{ + TNShardRecord: metadata.TNShardRecord{ShardID: 2}, + ReplicaID: 200, + }}, + }}, + refreshFails: test.refreshFails, + } + runtime.ServiceRuntime(s.sid).SetGlobalVariables(runtime.ClusterService, cluster) + + meta := NewTestTxn(1, 1, 1) + meta.Status = txn.TxnStatus_Prepared + meta.PreparedTS = NewTestTimestamp(2) + meta.TNShards = append(meta.TNShards, test.persistedRoute) + require.NoError(t, s.stopper.RunTask(s.doRecovery)) + s.txnC <- meta + close(s.txnC) + s.waitRecoveryCompleted() + + select { + case <-sender.notifyC: + case <-time.After(time.Second): + t.Fatal("recovery did not query the remote participant") + } + sender.mu.Lock() + require.NotEmpty(t, sender.requests) + assert.Equal(t, uint64(200), sender.requests[0][0].GetTargetTN().ReplicaID) + assert.Equal(t, "current-tn-2", sender.requests[0][0].GetTargetTN().Address) + sender.mu.Unlock() + + cluster.mu.Lock() + assert.Equal(t, test.expectRefreshes, cluster.refreshCalls) + cluster.mu.Unlock() + }) + } +} + +func TestRecoveryReusesAuthoritativeSnapshotAcrossTransactions(t *testing.T) { + sender := newRecoveryRouteSender() + s := NewTestTxnServiceWithLog(t, 1, sender, NewTestClock(0), nil).(*service) + defer s.stopper.Stop() + cluster := &staleRecoveryCluster{ + current: []metadata.TNService{{ + TxnServiceAddress: "current-tn-2", + Shards: []metadata.TNShard{{ + TNShardRecord: metadata.TNShardRecord{ShardID: 2}, + ReplicaID: 200, + }}, + }}, + } + runtime.ServiceRuntime(s.sid).SetGlobalVariables(runtime.ClusterService, cluster) + + require.NoError(t, s.stopper.RunTask(s.doRecovery)) + for txnID := byte(1); txnID <= 2; txnID++ { + meta := NewTestTxn(txnID, 1, 1) + meta.Status = txn.TxnStatus_Prepared + meta.PreparedTS = NewTestTimestamp(2) + meta.TNShards = append(meta.TNShards, metadata.TNShard{ + TNShardRecord: metadata.TNShardRecord{ShardID: 2}, + }) + s.txnC <- meta + } + close(s.txnC) + s.waitRecoveryCompleted() + + cluster.mu.Lock() + assert.Equal(t, 1, cluster.refreshCalls) + cluster.mu.Unlock() +} + +func TestRecoveryRestoresRoutesBeforeCommitTNShard(t *testing.T) { + sender := newRecoveryRouteSender() + sender.block = true + s := NewTestTxnServiceWithLog(t, 1, sender, NewTestClock(0), nil).(*service) + remote := metadata.TNShard{ + TNShardRecord: metadata.TNShardRecord{ShardID: 2}, + ReplicaID: 200, + } + installRecoveryCluster(t, s.sid, metadata.TNService{ + TxnServiceAddress: "tn-2", + Shards: []metadata.TNShard{remote}, + }) + + meta := NewTestTxn(1, 1, 1) + meta.Status = txn.TxnStatus_Committing + meta.PreparedTS = NewTestTimestamp(2) + meta.CommitTS = NewTestTimestamp(3) + meta.TNShards = append(meta.TNShards, metadata.TNShard{ + TNShardRecord: metadata.TNShardRecord{ShardID: 2}, + }) + + require.NoError(t, s.stopper.RunTask(s.doRecovery)) + s.txnC <- meta + close(s.txnC) + s.waitRecoveryCompleted() + + select { + case <-sender.notifyC: + case <-time.After(time.Second): + t.Fatal("recovery did not send the commit request") + } + sender.mu.Lock() + require.NotEmpty(t, sender.requests) + require.Len(t, sender.requests[0], 1) + assert.Equal(t, txn.TxnMethod_CommitTNShard, sender.requests[0][0].Method) + assert.Equal(t, metadata.TNShard{ + TNShardRecord: metadata.TNShardRecord{ShardID: 2}, + ReplicaID: 200, + Address: "tn-2", + }, sender.requests[0][0].GetTargetTN()) + sender.mu.Unlock() + + require.NoError(t, s.Close(false)) +} + +func TestRecoveryWaitsUntilParticipantRouteAppears(t *testing.T) { + sender := newRecoveryRouteSender() + s := NewTestTxnServiceWithLog(t, 1, sender, NewTestClock(0), nil).(*service) + defer s.stopper.Stop() + client := &recoveryClusterClient{} + c := clusterservice.NewMOCluster(s.sid, client, 50*time.Millisecond) + t.Cleanup(c.Close) + runtime.ServiceRuntime(s.sid).SetGlobalVariables(runtime.ClusterService, c) + + meta := NewTestTxn(1, 1, 1) + meta.Status = txn.TxnStatus_Prepared + meta.PreparedTS = NewTestTimestamp(2) + meta.TNShards = append(meta.TNShards, metadata.TNShard{ + TNShardRecord: metadata.TNShardRecord{ShardID: 2}, + }) + require.NoError(t, s.stopper.RunTask(s.doRecovery)) + s.txnC <- meta + close(s.txnC) + + select { + case <-sender.notifyC: + t.Fatal("recovery sent a request before the participant route was available") + case <-time.After(150 * time.Millisecond): + } + client.setTNRoute(2, 200, "tn-2") + s.waitRecoveryCompleted() + + select { + case <-sender.notifyC: + case <-time.After(time.Second): + t.Fatal("recovery did not resume after the participant route appeared") + } + sender.mu.Lock() + defer sender.mu.Unlock() + require.NotEmpty(t, sender.requests) + assert.Equal(t, "tn-2", sender.requests[0][0].GetTargetTN().Address) +} + +func TestRecoveryMissingParticipantRouteStopsWhenCanceled(t *testing.T) { + sender := newRecoveryRouteSender() + meta := NewTestTxn(1, 1, 1) + meta.Status = txn.TxnStatus_Prepared + meta.PreparedTS = NewTestTimestamp(2) + meta.TNShards = append(meta.TNShards, metadata.TNShard{ + TNShardRecord: metadata.TNShardRecord{ShardID: 99}, + }) + mlog := mem.NewMemLog() + addLog(t, mlog, meta, 1) + s := NewTestTxnServiceWithLog(t, 1, sender, NewTestClock(0), mlog).(*service) + installRecoveryCluster(t, s.sid) + + started := make(chan error, 1) + go func() { started <- s.Start() }() + select { + case <-s.recoveryC: + t.Fatal("recovery silently completed with an unresolved participant") + case <-time.After(100 * time.Millisecond): + } + + closed := make(chan error, 1) + go func() { closed <- s.Close(false) }() + select { + case err := <-closed: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("Close deadlocked while recovery waited for cluster metadata") + } + select { + case err := <-started: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("Start remained blocked after Close canceled recovery") + } + + sender.mu.Lock() + defer sender.mu.Unlock() + assert.Empty(t, sender.requests, "must not send to an unresolved or empty address") +} + +func TestRecoveryParticipantDoesNotWaitForUnrelatedRoute(t *testing.T) { + sender := newRecoveryRouteSender() + s := NewTestTxnServiceWithLog(t, 2, sender, NewTestClock(0), nil).(*service) + defer s.stopper.Stop() + defer s.CancelRecovery() + installRecoveryCluster(t, s.sid) + + meta := NewTestTxn(1, 1, 1, 2, 3) + meta.Status = txn.TxnStatus_Prepared + meta.PreparedTS = NewTestTimestamp(2) + require.NoError(t, s.stopper.RunTask(s.doRecovery)) + s.txnC <- meta + close(s.txnC) + + select { + case <-s.recoveryC: + case <-time.After(time.Second): + t.Fatal("non-coordinator recovery waited for an unrelated participant route") + } + assert.NotNil(t, s.getTxnContext(meta.ID)) + sender.mu.Lock() + defer sender.mu.Unlock() + assert.Empty(t, sender.requests) +} + +func TestRecoveryCommittedTerminalRecordDoesNotWaitForObsoleteRoute(t *testing.T) { + sender := newRecoveryRouteSender() + meta := NewTestTxn(1, 1, 1) + meta.Status = txn.TxnStatus_Prepared + meta.PreparedTS = NewTestTimestamp(2) + meta.TNShards = append(meta.TNShards, metadata.TNShard{ + TNShardRecord: metadata.TNShardRecord{ShardID: 99}, + }) + mlog := mem.NewMemLog() + addLog(t, mlog, meta, 1) + meta.Status = txn.TxnStatus_Committed + meta.CommitTS = NewTestTimestamp(3) + addLog(t, mlog, meta) + + s := NewTestTxnServiceWithLog(t, 1, sender, NewTestClock(0), mlog).(*service) + installRecoveryCluster(t, s.sid) + started := make(chan error, 1) + go func() { started <- s.Start() }() + + select { + case err := <-started: + require.NoError(t, err) + case <-time.After(time.Second): + s.CancelRecovery() + require.NoError(t, <-started) + t.Fatal("obsolete prepared route blocked a later committed terminal record") + } + defer func() { require.NoError(t, s.Close(false)) }() + assert.Nil(t, s.getTxnContext(meta.ID)) + sender.mu.Lock() + defer sender.mu.Unlock() + assert.Empty(t, sender.requests) +} + +func TestRecoveryRetriesTransientClusterLookupError(t *testing.T) { + sender := newRecoveryRouteSender() + s := NewTestTxnServiceWithLog(t, 1, sender, NewTestClock(0), nil).(*service) + defer s.stopper.Stop() + runtime.ServiceRuntime(s.sid).SetGlobalVariables(runtime.ClusterService, "not-a-cluster") + + meta := NewTestTxn(1, 1, 1) + meta.Status = txn.TxnStatus_Prepared + meta.PreparedTS = NewTestTimestamp(2) + meta.TNShards = append(meta.TNShards, metadata.TNShard{ + TNShardRecord: metadata.TNShardRecord{ShardID: 2}, + }) + require.NoError(t, s.stopper.RunTask(s.doRecovery)) + s.txnC <- meta + close(s.txnC) + select { + case <-s.recoveryC: + t.Fatal("recovery silently completed after a transient cluster lookup error") + case <-time.After(100 * time.Millisecond): + } + + installRecoveryCluster(t, s.sid, metadata.TNService{ + TxnServiceAddress: "tn-2", + Shards: []metadata.TNShard{{ + TNShardRecord: metadata.TNShardRecord{ShardID: 2}, + ReplicaID: 200, + }}, + }) + s.waitRecoveryCompleted() + select { + case <-sender.notifyC: + case <-time.After(time.Second): + t.Fatal("recovery did not resume after cluster lookup recovered") + } + sender.mu.Lock() + defer sender.mu.Unlock() + require.NotEmpty(t, sender.requests) + assert.Equal(t, "tn-2", sender.requests[0][0].GetTargetTN().Address) +} + func TestRecoveryFromCommittedWithData(t *testing.T) { mlog := mem.NewMemLog() wTxn := NewTestTxn(1, 1, 1) @@ -145,6 +651,7 @@ func TestRecoveryFromMultiTNShardWithAllPrepared(t *testing.T) { s1 := NewTestTxnServiceWithLog(t, 1, sender, NewTestClock(0), mlog1).(*service) s2 := NewTestTxnServiceWithLog(t, 2, sender, NewTestClock(0), mlog2).(*service) + installRecoveryRoutesForServices(t, s1, s2) sender.AddTxnService(s1) sender.AddTxnService(s2) @@ -191,6 +698,7 @@ func TestRecoveryFromMultiTNShardWithAnyNotPrepared(t *testing.T) { s1 := NewTestTxnServiceWithLog(t, 1, sender, NewTestClock(0), mlog1).(*service) s2 := NewTestTxnServiceWithLog(t, 2, sender, NewTestClock(0), mlog2).(*service) + installRecoveryRoutesForServices(t, s1, s2) sender.AddTxnService(s1) sender.AddTxnService(s2) @@ -236,6 +744,7 @@ func TestRecoveryFromMultiTNShardWithCommitting(t *testing.T) { s1 := NewTestTxnServiceWithLog(t, 1, sender, NewTestClock(0), mlog1).(*service) s2 := NewTestTxnServiceWithLog(t, 2, sender, NewTestClock(0), mlog2).(*service) + installRecoveryRoutesForServices(t, s1, s2) sender.AddTxnService(s1) sender.AddTxnService(s2) @@ -277,6 +786,10 @@ func TestRecoveryEndDoesNotPanicWhenStopperUnavailable(t *testing.T) { }() s := NewTestTxnServiceWithLog(t, 1, sender, NewTestClock(0), nil).(*service) + installRecoveryCluster(t, s.sid, metadata.TNService{ + TxnServiceAddress: "dn-2", + Shards: []metadata.TNShard{NewTestTNShard(2)}, + }) s.stopper.Stop() wTxn := NewTestTxn(1, 1, 1, 2) diff --git a/pkg/txn/service/types.go b/pkg/txn/service/types.go index 25f1e403ae266..331bf18b38f74 100644 --- a/pkg/txn/service/types.go +++ b/pkg/txn/service/types.go @@ -33,6 +33,8 @@ type TxnService interface { Shard() metadata.TNShard // Start start the txn service Start() error + // CancelRecovery interrupts a Start blocked in recovery without closing storage. + CancelRecovery() // Close close the txn service. Destroy TxnStorage if destroy is true. Close(destroy bool) error diff --git a/pkg/txn/storage/mem/kv_txn_storage.go b/pkg/txn/storage/mem/kv_txn_storage.go index ba6a79d6b34e1..8b153e0c75e58 100644 --- a/pkg/txn/storage/mem/kv_txn_storage.go +++ b/pkg/txn/storage/mem/kv_txn_storage.go @@ -143,6 +143,9 @@ func (kv *KVTxnStorage) StartRecovery(ctx context.Context, c chan txn.TxnMeta) { for { logs, lsn, err := kv.logClient.Read(ctx, kv.recoverFrom, math.MaxUint64) if err != nil { + if ctx.Err() != nil { + return + } panic(err) } @@ -177,7 +180,11 @@ func (kv *KVTxnStorage) StartRecovery(ctx context.Context, c chan txn.TxnMeta) { panic(fmt.Sprintf("invalid txn status %s", klog.Txn.Status.String())) } - c <- klog.Txn + select { + case <-ctx.Done(): + return + case c <- klog.Txn: + } } } diff --git a/pkg/txn/storage/mem/kv_txn_storage_test.go b/pkg/txn/storage/mem/kv_txn_storage_test.go index f5c5d686139c0..bd3453c6aea3a 100644 --- a/pkg/txn/storage/mem/kv_txn_storage_test.go +++ b/pkg/txn/storage/mem/kv_txn_storage_test.go @@ -20,10 +20,12 @@ import ( "math" "sync/atomic" "testing" + "time" "github.com/google/uuid" "github.com/matrixorigin/matrixone/pkg/common/moerr" "github.com/matrixorigin/matrixone/pkg/logservice" + logpb "github.com/matrixorigin/matrixone/pkg/pb/logservice" "github.com/matrixorigin/matrixone/pkg/pb/timestamp" "github.com/matrixorigin/matrixone/pkg/pb/txn" "github.com/matrixorigin/matrixone/pkg/txn/clock" @@ -35,6 +37,21 @@ type closeTrackingLogClient struct { closed atomic.Int32 } +type cancelBlockingLogClient struct { + logservice.Client + started chan struct{} +} + +func (c *cancelBlockingLogClient) Read( + ctx context.Context, + firstLsn logservice.Lsn, + _ uint64, +) ([]logpb.LogRecord, logservice.Lsn, error) { + close(c.started) + <-ctx.Done() + return nil, firstLsn, ctx.Err() +} + func (c *closeTrackingLogClient) Close() error { c.closed.Add(1) return c.Client.Close() @@ -299,6 +316,64 @@ func TestRecovery(t *testing.T) { assert.Equal(t, 6, len(txns)) } +func TestRecoveryCancellationUnblocksFullTxnChannel(t *testing.T) { + l := NewMemLog() + s := NewKVTxnStorage(0, l, newTestClock(1)) + recoveryC := make(chan txn.TxnMeta, 16) + for i := 0; i <= cap(recoveryC); i++ { + meta := txn.TxnMeta{ + ID: []byte{byte(i + 1)}, + Status: txn.TxnStatus_Committed, + } + _, err := s.saveLog(&KVLog{Txn: meta}) + assert.NoError(t, err) + } + + ctx, cancel := context.WithCancel(context.Background()) + recovered := make(chan struct{}) + go func() { + NewKVTxnStorage(1, l, newTestClock(1)).StartRecovery(ctx, recoveryC) + close(recovered) + }() + + assert.Eventually(t, func() bool { + return len(recoveryC) == cap(recoveryC) + }, time.Second, time.Millisecond) + cancel() + + select { + case <-recovered: + case <-time.After(time.Second): + t.Fatal("recovery remained blocked on a full transaction channel after cancellation") + } +} + +func TestRecoveryCancellationStopsLogRead(t *testing.T) { + client := &cancelBlockingLogClient{ + Client: NewMemLog(), + started: make(chan struct{}), + } + storage := NewKVTxnStorage(1, client, newTestClock(1)) + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + go func() { + storage.StartRecovery(ctx, make(chan txn.TxnMeta)) + close(done) + }() + + select { + case <-client.started: + case <-time.After(time.Second): + t.Fatal("recovery did not start the log read") + } + cancel() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("recovery log read did not stop after cancellation") + } +} + func TestEvent(t *testing.T) { l := NewMemLog() s := NewKVTxnStorage(0, l, newTestClock(1)) diff --git a/pkg/vm/engine/test/testutil/tae_engine.go b/pkg/vm/engine/test/testutil/tae_engine.go index f38cfca0e2727..9f2f70b31621f 100644 --- a/pkg/vm/engine/test/testutil/tae_engine.go +++ b/pkg/vm/engine/test/testutil/tae_engine.go @@ -63,7 +63,8 @@ func (ts *TestTxnStorage) Shard() metadata.TNShard { return GetDefaultTNShard() } -func (ts *TestTxnStorage) Start() error { return nil } +func (ts *TestTxnStorage) Start() error { return nil } +func (ts *TestTxnStorage) CancelRecovery() {} func (ts *TestTxnStorage) Close(destroy bool) error { var firstErr error if err := ts.GetDB().Close(); err != nil {