From de43121d9cfc74000e0a4f0fa442e40027d46e58 Mon Sep 17 00:00:00 2001 From: aptend Date: Wed, 15 Jul 2026 14:20:12 +0800 Subject: [PATCH] fix: complete TN lifecycle hardening after replica changes --- pkg/tnservice/cfg.go | 5 + pkg/tnservice/cfg_test.go | 9 ++ pkg/tnservice/factory.go | 24 +--- pkg/tnservice/factory_test.go | 24 ++++ pkg/tnservice/replica.go | 132 ++++++++++++++++-- pkg/tnservice/replica_test.go | 99 +++++++++++++ pkg/tnservice/store.go | 30 ++-- pkg/tnservice/store_rpc_handler.go | 84 ++++++----- pkg/tnservice/store_rpc_handler_test.go | 99 +++++++++++-- pkg/tnservice/store_test.go | 90 +++++++++++- pkg/txn/service/service.go | 18 +-- pkg/txn/service/service_cn_handler.go | 24 +++- pkg/txn/service/service_cn_handler_test.go | 99 +++++++++++++ pkg/txn/service/service_test.go | 22 +++ pkg/txn/storage/mem/kv_txn_storage.go | 11 +- pkg/txn/storage/mem/kv_txn_storage_test.go | 19 +++ pkg/txn/storage/tae/storage.go | 3 + pkg/txn/storage/tae/storage_lifecycle_test.go | 28 +++- pkg/vm/engine/tae/logtail/service/response.go | 47 +++++-- .../tae/logtail/service/response_test.go | 7 + pkg/vm/engine/tae/logtail/service/server.go | 6 +- .../engine/tae/logtail/service/server_test.go | 24 ++++ 22 files changed, 781 insertions(+), 123 deletions(-) diff --git a/pkg/tnservice/cfg.go b/pkg/tnservice/cfg.go index a715a7b9125f4..6fbbaaa66e20e 100644 --- a/pkg/tnservice/cfg.go +++ b/pkg/tnservice/cfg.go @@ -31,6 +31,7 @@ import ( "github.com/matrixorigin/matrixone/pkg/pb/txn" "github.com/matrixorigin/matrixone/pkg/txn/rpc" "github.com/matrixorigin/matrixone/pkg/util/toml" + logtailservice "github.com/matrixorigin/matrixone/pkg/vm/engine/tae/logtail/service" "github.com/matrixorigin/matrixone/pkg/vm/engine/tae/options" ) @@ -282,6 +283,10 @@ func (c *Config) Validate() error { if c.LogtailServer.RpcMaxMessageSize <= 0 { c.LogtailServer.RpcMaxMessageSize = toml.ByteSize(defaultRpcMaxMsgSize) } + if err := logtailservice.ValidateRPCMaxMessageSize( + int64(c.LogtailServer.RpcMaxMessageSize)); err != nil { + return err + } if c.LogtailServer.LogtailRPCStreamPoisonTime.Duration <= 0 { c.LogtailServer.LogtailRPCStreamPoisonTime.Duration = defaultRPCStreamPoisonTime } diff --git a/pkg/tnservice/cfg_test.go b/pkg/tnservice/cfg_test.go index 094b3fdb01420..a251d07c49c19 100644 --- a/pkg/tnservice/cfg_test.go +++ b/pkg/tnservice/cfg_test.go @@ -17,10 +17,19 @@ package tnservice import ( "testing" + "github.com/matrixorigin/matrixone/pkg/util/toml" "github.com/matrixorigin/matrixone/pkg/vm/engine/tae/options" "github.com/stretchr/testify/assert" ) +func TestValidateRejectsInvalidLogtailRPCMessageSize(t *testing.T) { + for _, size := range []toml.ByteSize{1, 101 * 1024 * 1024} { + c := &Config{UUID: "tn1"} + c.LogtailServer.RpcMaxMessageSize = size + assert.Error(t, c.Validate()) + } +} + func TestValidate(t *testing.T) { c := &Config{} assert.Error(t, c.Validate()) diff --git a/pkg/tnservice/factory.go b/pkg/tnservice/factory.go index 9a9d56b39c168..72d3cf9ec2116 100644 --- a/pkg/tnservice/factory.go +++ b/pkg/tnservice/factory.go @@ -31,7 +31,6 @@ import ( "github.com/matrixorigin/matrixone/pkg/vm/engine/memoryengine" "github.com/matrixorigin/matrixone/pkg/vm/engine/tae/logstore/driver/logservicedriver" "github.com/matrixorigin/matrixone/pkg/vm/engine/tae/options" - "go.uber.org/zap" ) var ( @@ -47,29 +46,12 @@ func (s *store) createTxnStorage( shard metadata.TNShard, txnServer rpc.TxnServer, ) (storage.TxnStorage, error) { - - factory := s.createLogServiceClientFactroy(shard) - closeLogClientFn := func(logClient logservice.Client) { - if err := logClient.Close(); err != nil { - s.rt.Logger().Error("close log client failed", - zap.Error(err)) - } - } - switch s.cfg.Txn.Storage.Backend { case StorageMEM: - logClient, err := factory() - if err != nil { - return nil, err - } - ts, err := s.newMemTxnStorage(shard, logClient, s.hakeeperClient) - if err != nil { - closeLogClientFn(logClient) - return nil, err - } - return ts, nil + return s.newMemTxnStorage(shard, s.hakeeperClient) case StorageMEMKV: + factory := s.createLogServiceClientFactroy(shard) logClient, err := factory() if err != nil { return nil, err @@ -77,6 +59,7 @@ func (s *store) createTxnStorage( return s.newMemKVStorage(shard, logClient) case StorageTAE: + factory := s.createLogServiceClientFactroy(shard) ts, err := s.newTAEStorage(ctx, shard, factory, txnServer) if err != nil { return nil, err @@ -119,7 +102,6 @@ func (s *store) newLogServiceClient(shard metadata.TNShard) (logservice.Client, func (s *store) newMemTxnStorage( shard metadata.TNShard, - logClient logservice.Client, hakeeper logservice.TNHAKeeperClient, ) (storage.TxnStorage, error) { // should it be no fixed or a certain size? diff --git a/pkg/tnservice/factory_test.go b/pkg/tnservice/factory_test.go index 8c575781c3d93..4c0d66b6ace7c 100644 --- a/pkg/tnservice/factory_test.go +++ b/pkg/tnservice/factory_test.go @@ -18,6 +18,7 @@ import ( "context" "testing" + "github.com/matrixorigin/matrixone/pkg/catalog" "github.com/matrixorigin/matrixone/pkg/common/runtime" "github.com/matrixorigin/matrixone/pkg/common/stopper" "github.com/matrixorigin/matrixone/pkg/logservice" @@ -53,3 +54,26 @@ func TestCreateTxnStorage(t *testing.T) { assert.Error(t, err) assert.Nil(t, v) } + +func TestCreateMemoryStorageDoesNotCreateUnusedLogClient(t *testing.T) { + catalog.SetupDefines("") + ctx := context.Background() + s := &store{ + rt: runtime.DefaultRuntime(), + cfg: &Config{}, + stopper: stopper.NewStopper(""), + hakeeperClient: newTestHAKeeperClient(), + } + calls := 0 + s.options.logServiceClientFactory = func(metadata.TNShard) (logservice.Client, error) { + calls++ + return mem.NewMemLog(), nil + } + s.cfg.Txn.Storage.Backend = StorageMEM + + v, err := s.createTxnStorage(ctx, metadata.TNShard{}, nil) + assert.NoError(t, err) + assert.NotNil(t, v) + assert.Zero(t, calls) + assert.NoError(t, v.Close(ctx)) +} diff --git a/pkg/tnservice/replica.go b/pkg/tnservice/replica.go index 7489325bbdf99..21b0f5264ac00 100644 --- a/pkg/tnservice/replica.go +++ b/pkg/tnservice/replica.go @@ -39,19 +39,45 @@ type replica struct { createCtx context.Context cancelCreate context.CancelFunc startedOnce sync.Once + closeOnce sync.Once + closeErr error mu struct { sync.RWMutex + cond *sync.Cond starting bool cancelled bool destroyOnCancel bool startErr error + activeCalls int } } +type txnServiceLease struct { + replica *replica + service service.TxnService + ctx context.Context + cancel context.CancelFunc + stopCancelPropagation func() bool + releaseOnce sync.Once +} + +func (l *txnServiceLease) release() { + l.releaseOnce.Do(func() { + l.stopCancelPropagation() + l.cancel() + l.replica.mu.Lock() + l.replica.mu.activeCalls-- + if l.replica.mu.activeCalls == 0 { + l.replica.mu.cond.Broadcast() + } + l.replica.mu.Unlock() + }) +} + func newReplica(shard metadata.TNShard, rt runtime.Runtime) *replica { ctx, cancel := context.WithCancel(context.Background()) - return &replica{ + r := &replica{ rt: rt, shard: shard, logger: rt.Logger().With(util.TxnTNShardField(shard)), @@ -59,6 +85,8 @@ func newReplica(shard metadata.TNShard, rt runtime.Runtime) *replica { createCtx: ctx, cancelCreate: cancel, } + r.mu.cond = sync.NewCond(&r.mu.RWMutex) + return r } func (r *replica) start(txnService service.TxnService) error { @@ -124,49 +152,131 @@ func (r *replica) waitStartCompleted() { } func (r *replica) close(destroy bool) error { + r.cancelStart(destroy) + r.closeOnce.Do(func() { + r.closeErr = r.closeOnceFn() + }) + return r.closeErr +} + +func (r *replica) closeOnceFn() error { r.mu.RLock() starting := r.mu.starting r.mu.RUnlock() if !starting { return nil } - startErr := r.waitStarted(context.Background()) - r.mu.RLock() + + r.waitStartCompleted() + r.mu.Lock() + for r.mu.activeCalls > 0 { + r.mu.cond.Wait() + } + startErr := r.mu.startErr txnService := r.service - r.mu.RUnlock() + destroy := r.mu.destroyOnCancel + r.mu.Unlock() if txnService == nil { return startErr } return errors.Join(startErr, txnService.Close(destroy)) } +func (r *replica) started() bool { + select { + case <-r.startedC: + r.mu.RLock() + defer r.mu.RUnlock() + return !r.mu.cancelled && r.mu.startErr == nil && r.service != nil + default: + return false + } +} + func (r *replica) handleLocalRequest(ctx context.Context, request *txn.TxnRequest, response *txn.TxnResponse) error { - if err := r.waitStarted(ctx); err != nil { + lease, err := r.acquireService(ctx) + if err != nil { return err } + defer lease.release() prepareResponse(request, response) switch request.Method { case txn.TxnMethod_GetStatus: - return r.service.GetStatus(ctx, request, response) + return lease.service.GetStatus(lease.ctx, request, response) case txn.TxnMethod_Prepare: - return r.service.Prepare(ctx, request, response) + return lease.service.Prepare(lease.ctx, request, response) case txn.TxnMethod_CommitTNShard: - return r.service.CommitTNShard(ctx, request, response) + return lease.service.CommitTNShard(lease.ctx, request, response) case txn.TxnMethod_RollbackTNShard: - return r.service.RollbackTNShard(ctx, request, response) + return lease.service.RollbackTNShard(lease.ctx, request, response) default: - return moerr.NewNotSupportedf(ctx, "unknown txn request method: %s", request.Method.String()) + return moerr.NewNotSupportedf(lease.ctx, "unknown txn request method: %s", request.Method.String()) } } func (r *replica) waitStarted(ctx context.Context) error { + if err := context.Cause(ctx); err != nil { + return err + } + if err := context.Cause(r.createCtx); err != nil { + return err + } select { case <-ctx.Done(): - return ctx.Err() + return context.Cause(ctx) + case <-r.createCtx.Done(): + return context.Cause(r.createCtx) case <-r.startedC: + if err := context.Cause(ctx); err != nil { + return err + } r.mu.RLock() defer r.mu.RUnlock() + if r.mu.cancelled { + return context.Canceled + } return r.mu.startErr } } + +// acquireService pins the transaction service until release is called. Closing +// a replica cancels the call context and waits for all acquired services to be +// released before closing the underlying storage. +func (r *replica) acquireService(ctx context.Context) (*txnServiceLease, error) { + if err := r.waitStarted(ctx); err != nil { + return nil, err + } + + r.mu.Lock() + if err := context.Cause(ctx); err != nil { + r.mu.Unlock() + return nil, err + } + if r.mu.cancelled { + r.mu.Unlock() + return nil, context.Canceled + } + if r.mu.startErr != nil { + err := r.mu.startErr + r.mu.Unlock() + return nil, err + } + if r.service == nil { + r.mu.Unlock() + return nil, context.Canceled + } + txnService := r.service + r.mu.activeCalls++ + r.mu.Unlock() + + callCtx, cancel := context.WithCancel(ctx) + stopCancelPropagation := context.AfterFunc(r.createCtx, cancel) + return &txnServiceLease{ + replica: r, + service: txnService, + ctx: callCtx, + cancel: cancel, + stopCancelPropagation: stopCancelPropagation, + }, nil +} diff --git a/pkg/tnservice/replica_test.go b/pkg/tnservice/replica_test.go index ef8832b0cf1d1..7b4b471f2c23c 100644 --- a/pkg/tnservice/replica_test.go +++ b/pkg/tnservice/replica_test.go @@ -28,6 +28,8 @@ import ( "github.com/stretchr/testify/require" ) +const replicaTestTimeout = 30 * time.Second + type startErrorStorage struct { storage.TxnStorage startErr error @@ -37,6 +39,16 @@ type startErrorStorage struct { destroyCalls int } +type closeTrackingTxnService struct { + service.TxnService + closeCalls int +} + +func (s *closeTrackingTxnService) Close(destroy bool) error { + s.closeCalls++ + return s.TxnService.Close(destroy) +} + func (s *startErrorStorage) Start() error { return s.startErr } @@ -65,6 +77,93 @@ func TestCloseNotStartedReplica(t *testing.T) { assert.NoError(t, r.close(false)) } +func TestCloseStartedReplicaIsIdempotent(t *testing.T) { + r := newReplica(newTestTNShard(1, 2, 3), runtime.DefaultRuntime()) + sender := service.NewTestSender() + t.Cleanup(func() { require.NoError(t, sender.Close()) }) + base := service.NewTestTxnService(t, 1, sender, service.NewTestClock(1)) + txnService := &closeTrackingTxnService{TxnService: base} + require.NoError(t, r.start(txnService)) + + require.NoError(t, r.close(false)) + require.NoError(t, r.close(true)) + require.Equal(t, 1, txnService.closeCalls) +} + +func TestCloseStartedReplicaCancelsAndDrainsActiveCalls(t *testing.T) { + r := newReplica(newTestTNShard(1, 2, 3), runtime.DefaultRuntime()) + sender := service.NewTestSender() + t.Cleanup(func() { require.NoError(t, sender.Close()) }) + txnService := &closeTrackingTxnService{ + TxnService: service.NewTestTxnService(t, 1, sender, service.NewTestClock(1)), + } + require.NoError(t, r.start(txnService)) + + lease, err := r.acquireService(context.Background()) + require.NoError(t, err) + closed := make(chan error, 1) + go func() { + closed <- r.close(false) + }() + + select { + case <-lease.ctx.Done(): + require.ErrorIs(t, context.Cause(lease.ctx), context.Canceled) + case <-time.After(replicaTestTimeout): + t.Fatal("active call context was not canceled") + } + select { + case err := <-closed: + t.Fatalf("replica closed before active call was released: %v", err) + default: + } + nestedDone := make(chan error, 1) + go func() { + _, err := r.acquireService(context.Background()) + nestedDone <- err + }() + select { + case err := <-nestedDone: + require.ErrorIs(t, err, context.Canceled) + case <-time.After(replicaTestTimeout): + t.Fatal("nested local acquire blocked while close was draining") + } + + lease.release() + lease.release() + select { + case err := <-closed: + require.NoError(t, err) + case <-time.After(replicaTestTimeout): + t.Fatal("replica close did not drain") + } + require.Equal(t, 1, txnService.closeCalls) + + _, err = r.acquireService(context.Background()) + require.ErrorIs(t, err, context.Canceled) +} + +func TestAcquireServiceRejectsCanceledCaller(t *testing.T) { + r := newReplica(newTestTNShard(1, 2, 3), runtime.DefaultRuntime()) + sender := service.NewTestSender() + t.Cleanup(func() { require.NoError(t, sender.Close()) }) + txnService := service.NewTestTxnService(t, 1, sender, service.NewTestClock(1)) + require.NoError(t, r.start(txnService)) + t.Cleanup(func() { require.NoError(t, r.close(false)) }) + + cause := errors.New("caller canceled") + ctx, cancel := context.WithCancelCause(context.Background()) + cancel(cause) + for range 100 { + lease, err := r.acquireService(ctx) + require.Nil(t, lease) + require.ErrorIs(t, err, cause) + } + r.mu.RLock() + defer r.mu.RUnlock() + require.Zero(t, r.mu.activeCalls) +} + func TestCloseFailedStartReplica(t *testing.T) { startErr := errors.New("storage start failed") runtime.SetupServiceBasedRuntime("test", runtime.DefaultRuntime()) diff --git a/pkg/tnservice/store.go b/pkg/tnservice/store.go index 677d7a3d83f71..f376e37a514af 100644 --- a/pkg/tnservice/store.go +++ b/pkg/tnservice/store.go @@ -242,23 +242,32 @@ func (s *store) Close() error { 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. + s.replicas.Range(func(_, value any) bool { + value.(*replica).cancelStart(false) + return true + }) + if s.queryService != nil { + err = errors.Join(err, s.queryService.Close()) + } if s.cfg.ShardService.Enable { - err = s.shardServer.Close() + err = errors.Join(err, s.shardServer.Close()) } - err = errors.Join( - s.hakeeperClient.Close(), - s.sender.Close(), - s.server.Close(), - s.lockTableAllocator.Close(), - ) + err = errors.Join(err, s.server.Close()) s.replicas.Range(func(_, value any) bool { r := value.(*replica) - r.cancelStart(false) if e := r.close(false); e != nil { - err = errors.Join(e, err) + err = errors.Join(err, e) } return true }) + err = errors.Join( + err, + s.hakeeperClient.Close(), + s.sender.Close(), + s.lockTableAllocator.Close(), + ) if s.queryClient != nil { err = errors.Join(err, s.queryClient.Close()) } @@ -301,6 +310,9 @@ func (s *store) getTNShardInfo() []logservicepb.TNShardInfo { var shards []logservicepb.TNShardInfo s.replicas.Range(func(_, value any) bool { r := value.(*replica) + if !r.started() { + return true + } shards = append(shards, logservicepb.TNShardInfo{ ShardID: r.shard.ShardID, ReplicaID: r.shard.ReplicaID, diff --git a/pkg/tnservice/store_rpc_handler.go b/pkg/tnservice/store_rpc_handler.go index 02f630708ac02..5014c71d3e8a9 100644 --- a/pkg/tnservice/store_rpc_handler.go +++ b/pkg/tnservice/store_rpc_handler.go @@ -69,30 +69,33 @@ func (s *store) handleDebug(ctx context.Context, request *txn.TxnRequest, respon } func (s *store) doRead(ctx context.Context, request *txn.TxnRequest, response *txn.TxnResponse) error { - r, err := s.startedTNReplica(ctx, request, response) - if err != nil || r == nil { + lease, err := s.acquireTNReplica(ctx, request, response) + if err != nil || lease == nil { return err } + defer lease.release() prepareResponse(request, response) - return r.service.Read(ctx, request, response) + return lease.service.Read(lease.ctx, request, response) } func (s *store) doWrite(ctx context.Context, request *txn.TxnRequest, response *txn.TxnResponse) error { - r, err := s.startedTNReplica(ctx, request, response) - if err != nil || r == nil { + lease, err := s.acquireTNReplica(ctx, request, response) + if err != nil || lease == nil { return err } + defer lease.release() prepareResponse(request, response) - return r.service.Write(ctx, request, response) + return lease.service.Write(lease.ctx, request, response) } func (s *store) doDebug(ctx context.Context, request *txn.TxnRequest, response *txn.TxnResponse) error { - r, err := s.startedTNReplica(ctx, request, response) - if err != nil || r == nil { + lease, err := s.acquireTNReplica(ctx, request, response) + if err != nil || lease == nil { return err } + defer lease.release() prepareResponse(request, response) - return r.service.Debug(ctx, request, response) + return lease.service.Debug(lease.ctx, request, response) } func (s *store) handleCommit(ctx context.Context, request *txn.TxnRequest, response *txn.TxnResponse) error { @@ -100,57 +103,63 @@ func (s *store) handleCommit(ctx context.Context, request *txn.TxnRequest, respo trace.WithKind(trace.SpanKindStatement)) defer span.End() - r, err := s.startedTNReplica(ctx, request, response) - if err != nil || r == nil { + lease, err := s.acquireTNReplica(ctx, request, response) + if err != nil || lease == nil { return err } + defer lease.release() prepareResponse(request, response) - return r.service.Commit(ctx, request, response) + return lease.service.Commit(lease.ctx, request, response) } func (s *store) handleRollback(ctx context.Context, request *txn.TxnRequest, response *txn.TxnResponse) error { - r, err := s.startedTNReplica(ctx, request, response) - if err != nil || r == nil { + lease, err := s.acquireTNReplica(ctx, request, response) + if err != nil || lease == nil { return err } + defer lease.release() prepareResponse(request, response) - return r.service.Rollback(ctx, request, response) + return lease.service.Rollback(lease.ctx, request, response) } func (s *store) handlePrepare(ctx context.Context, request *txn.TxnRequest, response *txn.TxnResponse) error { - r, err := s.startedTNReplica(ctx, request, response) - if err != nil || r == nil { + lease, err := s.acquireTNReplica(ctx, request, response) + if err != nil || lease == nil { return err } + defer lease.release() prepareResponse(request, response) - return r.service.Prepare(ctx, request, response) + return lease.service.Prepare(lease.ctx, request, response) } func (s *store) handleCommitTNShard(ctx context.Context, request *txn.TxnRequest, response *txn.TxnResponse) error { - r, err := s.startedTNReplica(ctx, request, response) - if err != nil || r == nil { + lease, err := s.acquireTNReplica(ctx, request, response) + if err != nil || lease == nil { return err } + defer lease.release() prepareResponse(request, response) - return r.service.CommitTNShard(ctx, request, response) + return lease.service.CommitTNShard(lease.ctx, request, response) } func (s *store) handleRollbackTNShard(ctx context.Context, request *txn.TxnRequest, response *txn.TxnResponse) error { - r, err := s.startedTNReplica(ctx, request, response) - if err != nil || r == nil { + lease, err := s.acquireTNReplica(ctx, request, response) + if err != nil || lease == nil { return err } + defer lease.release() prepareResponse(request, response) - return r.service.RollbackTNShard(ctx, request, response) + return lease.service.RollbackTNShard(lease.ctx, request, response) } func (s *store) handleGetStatus(ctx context.Context, request *txn.TxnRequest, response *txn.TxnResponse) error { - r, err := s.startedTNReplica(ctx, request, response) - if err != nil || r == nil { + lease, err := s.acquireTNReplica(ctx, request, response) + if err != nil || lease == nil { return err } + defer lease.release() prepareResponse(request, response) - return r.service.GetStatus(ctx, request, response) + return lease.service.GetStatus(lease.ctx, request, response) } func (s *store) validTNShard(ctx context.Context, request *txn.TxnRequest, response *txn.TxnResponse) *replica { @@ -164,18 +173,19 @@ func (s *store) validTNShard(ctx context.Context, request *txn.TxnRequest, respo return r } -func (s *store) startedTNReplica( +func (s *store) acquireTNReplica( ctx context.Context, request *txn.TxnRequest, response *txn.TxnResponse, -) (*replica, error) { +) (*txnServiceLease, error) { r := s.validTNShard(ctx, request, response) if r == nil { return nil, nil } - if err := r.waitStarted(ctx); err != nil { + lease, err := r.acquireService(ctx) + if err != nil { if ctx.Err() != nil { - return nil, ctx.Err() + return nil, context.Cause(ctx) } response.TxnError = txn.WrapError( moerr.NewTNShardNotFound(ctx, s.cfg.UUID, r.shard.ShardID), @@ -183,7 +193,7 @@ func (s *store) startedTNReplica( ) return nil, nil } - return r, nil + return lease, nil } func prepareResponse(request *txn.TxnRequest, response *txn.TxnResponse) { @@ -230,8 +240,14 @@ func (s *store) maybeRetry(ctx context.Context, request *txn.TxnRequest, respons if wait == 0 { wait = defaultRetryInterval } - time.Sleep(wait) - return true + timer := time.NewTimer(wait) + defer timer.Stop() + select { + case <-ctx.Done(): + return false + case <-timer.C: + return context.Cause(ctx) == nil + } } } return false diff --git a/pkg/tnservice/store_rpc_handler_test.go b/pkg/tnservice/store_rpc_handler_test.go index 00cf8bf94e2b5..e5a41f03afe96 100644 --- a/pkg/tnservice/store_rpc_handler_test.go +++ b/pkg/tnservice/store_rpc_handler_test.go @@ -18,7 +18,9 @@ import ( "context" "errors" "sync" + "sync/atomic" "testing" + "testing/synctest" "time" "github.com/matrixorigin/matrixone/pkg/common/moerr" @@ -29,6 +31,64 @@ import ( "github.com/stretchr/testify/require" ) +func TestRetryWaitHonorsContext(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + req := txn.TxnRequest{Options: &txn.TxnRequestOptions{ + RetryCodes: []int32{int32(moerr.ErrTNShardNotFound)}, + RetryInterval: int64(time.Hour), + }} + resp := &txn.TxnResponse{} + var calls atomic.Int32 + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + var resultErr error + done := false + go func() { + resultErr = (&store{}).handleWithRetry(ctx, &req, resp, + func(context.Context, *txn.TxnRequest, *txn.TxnResponse) error { + calls.Add(1) + resp.TxnError = txn.WrapError(moerr.NewTNShardNotFoundNoCtx("tn", 1), 0) + return nil + }) + done = true + }() + + // The fake clock cannot advance while Wait is active, so returning here + // proves the request is blocked in the one-hour retry wait. + synctest.Wait() + require.False(t, done) + require.Equal(t, int32(1), calls.Load()) + + cancel() + synctest.Wait() + require.True(t, done) + require.NoError(t, resultErr) + require.Equal(t, int32(1), calls.Load()) + require.Equal(t, uint32(moerr.ErrTNShardNotFound), resp.TxnError.Code) + }) +} + +func TestRetryDoesNotRedispatchWhenContextAndTimerAreReady(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + req := txn.TxnRequest{Options: &txn.TxnRequestOptions{ + RetryCodes: []int32{int32(moerr.ErrTNShardNotFound)}, + RetryInterval: -1, + }} + resp := &txn.TxnResponse{} + var calls atomic.Int32 + + err := (&store{}).handleWithRetry(ctx, &req, resp, + func(context.Context, *txn.TxnRequest, *txn.TxnResponse) error { + calls.Add(1) + resp.TxnError = txn.WrapError(moerr.NewTNShardNotFoundNoCtx("tn", 1), 0) + cancel() + return nil + }) + + assert.NoError(t, err) + assert.Equal(t, int32(1), calls.Load()) +} + func TestStartedTNReplicaLifecycle(t *testing.T) { newStore := func() *store { return &store{cfg: &Config{UUID: "test"}, replicas: &sync.Map{}} @@ -41,6 +101,9 @@ func TestStartedTNReplicaLifecycle(t *testing.T) { newReadyReplica := func(t *testing.T, shardID, replicaID uint64) *replica { r := newReplica(newTestTNShard(shardID, replicaID, shardID), runtime.DefaultRuntime()) require.True(t, r.reserveStart()) + r.mu.Lock() + r.service = &closeTrackingTxnService{} + r.mu.Unlock() r.finishStart(nil) return r } @@ -55,9 +118,11 @@ func TestStartedTNReplicaLifecycle(t *testing.T) { s.replicas.Store(uint64(1), r) req := newRequest(1, 2) - got, err := s.startedTNReplica(context.Background(), &req, &txn.TxnResponse{}) + got, err := s.acquireTNReplica(context.Background(), &req, &txn.TxnResponse{}) require.NoError(t, err) - require.Same(t, r, got) + require.NotNil(t, got) + defer got.release() + require.Same(t, r.service, got.service) }) t.Run("wait then ready", func(t *testing.T) { @@ -67,18 +132,24 @@ func TestStartedTNReplicaLifecycle(t *testing.T) { s.replicas.Store(uint64(1), r) req := newRequest(1, 2) entered := make(chan struct{}) - result := make(chan *replica, 1) + result := make(chan *txnServiceLease, 1) errC := make(chan error, 1) go func() { close(entered) - got, err := s.startedTNReplica(context.Background(), &req, &txn.TxnResponse{}) + got, err := s.acquireTNReplica(context.Background(), &req, &txn.TxnResponse{}) result <- got errC <- err }() <-entered + r.mu.Lock() + r.service = &closeTrackingTxnService{} + r.mu.Unlock() r.finishStart(nil) - require.Same(t, r, <-result) + lease := <-result + require.NotNil(t, lease) + defer lease.release() + require.Same(t, r.service, lease.service) require.NoError(t, <-errC) }) @@ -93,7 +164,7 @@ func TestStartedTNReplicaLifecycle(t *testing.T) { errC := make(chan error, 1) go func() { close(entered) - _, err := s.startedTNReplica(ctx, &req, &txn.TxnResponse{}) + _, err := s.acquireTNReplica(ctx, &req, &txn.TxnResponse{}) errC <- err }() @@ -118,7 +189,7 @@ func TestStartedTNReplicaLifecycle(t *testing.T) { req := newRequest(1, 2) response := &txn.TxnResponse{} - got, err := s.startedTNReplica(context.Background(), &req, response) + got, err := s.acquireTNReplica(context.Background(), &req, response) require.NoError(t, err) require.Nil(t, got) requireNotFound(t, response) @@ -130,23 +201,27 @@ func TestStartedTNReplicaLifecycle(t *testing.T) { oldReplica := newReadyReplica(t, 1, 2) s.replicas.Store(uint64(1), oldReplica) oldRequest := newRequest(1, 2) - got, err := s.startedTNReplica(context.Background(), &oldRequest, &txn.TxnResponse{}) + got, err := s.acquireTNReplica(context.Background(), &oldRequest, &txn.TxnResponse{}) require.NoError(t, err) - require.Same(t, oldReplica, got) + require.NotNil(t, got) + require.Same(t, oldReplica.service, got.service) + got.release() s.replicas.Delete(uint64(1)) replacement := newReadyReplica(t, 1, 4) s.replicas.Store(uint64(1), replacement) oldResponse := &txn.TxnResponse{} - got, err = s.startedTNReplica(context.Background(), &oldRequest, oldResponse) + got, err = s.acquireTNReplica(context.Background(), &oldRequest, oldResponse) require.NoError(t, err) require.Nil(t, got) requireNotFound(t, oldResponse) newRequest := newRequest(1, 4) - got, err = s.startedTNReplica(context.Background(), &newRequest, &txn.TxnResponse{}) + got, err = s.acquireTNReplica(context.Background(), &newRequest, &txn.TxnResponse{}) require.NoError(t, err) - require.Same(t, replacement, got) + require.NotNil(t, got) + defer got.release() + require.Same(t, replacement.service, got.service) }) } diff --git a/pkg/tnservice/store_test.go b/pkg/tnservice/store_test.go index bf344b2537dd5..360bf25199aac 100644 --- a/pkg/tnservice/store_test.go +++ b/pkg/tnservice/store_test.go @@ -31,8 +31,10 @@ 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/queryservice" "github.com/matrixorigin/matrixone/pkg/queryservice/client" "github.com/matrixorigin/matrixone/pkg/txn/clock" + "github.com/matrixorigin/matrixone/pkg/txn/rpc" "github.com/matrixorigin/matrixone/pkg/txn/service" "github.com/matrixorigin/matrixone/pkg/txn/storage/mem" "github.com/stretchr/testify/assert" @@ -44,11 +46,33 @@ var ( testTNLogtailAddress = "127.0.0.1:22001" ) +const storeTestTimeout = 30 * time.Second + type storeQueryClient struct { client.QueryClient closeCalls int } +type storeQueryService struct { + queryservice.QueryService + closeCalls int +} + +type storeTxnServer struct { + rpc.TxnServer + beforeClose func() +} + +func (s *storeTxnServer) Close() error { + s.beforeClose() + return s.TxnServer.Close() +} + +func (s *storeQueryService) Close() error { + s.closeCalls++ + return s.QueryService.Close() +} + func (c *storeQueryClient) Close() error { c.closeCalls++ return nil @@ -255,14 +279,60 @@ func TestStoreCloseClosesSharedQueryClient(t *testing.T) { }) } -func TestReplicaCreateRetryStopsPromptly(t *testing.T) { +func TestStoreCloseClosesQueryService(t *testing.T) { + var tracked *storeQueryService + runTNStoreTest(t, func(s *store) { + tracked = &storeQueryService{QueryService: s.queryService} + s.queryService = tracked + t.Cleanup(func() { + require.Equal(t, 1, tracked.closeCalls) + }) + }) +} + +func TestStoreCloseCancelsReplicasBeforeDrainingRPCServer(t *testing.T) { + var canceledBeforeServerClose atomic.Bool + runTNStoreTest(t, func(s *store) { + r := newReplica(newTestTNShard(1, 2, 3), s.rt) + s.replicas.Store(r.shard.ShardID, r) + s.server = &storeTxnServer{ + TxnServer: s.server, + beforeClose: func() { + r.mu.RLock() + defer r.mu.RUnlock() + canceledBeforeServerClose.Store(r.mu.cancelled) + }, + } + }) + require.True(t, canceledBeforeServerClose.Load()) +} + +func TestHeartbeatOnlyReportsStartedReplicas(t *testing.T) { + runTNStoreTest(t, func(s *store) { + r := newReplica(newTestTNShard(1, 2, 3), s.rt) + s.replicas.Store(r.shard.ShardID, r) + require.Empty(t, s.getTNShardInfo()) + + sender := service.NewTestSender() + t.Cleanup(func() { require.NoError(t, sender.Close()) }) + txnService := service.NewTestTxnService(t, 1, sender, service.NewTestClock(1)) + require.NoError(t, r.start(txnService)) + require.Equal(t, []logservicepb.TNShardInfo{{ + ShardID: 1, ReplicaID: 2, + }}, s.getTNShardInfo()) + }) +} + +func TestReplicaCreateRetryStopsOnStoreStop(t *testing.T) { oldInterval := retryCreateStorageInterval - retryCreateStorageInterval = time.Second + retryCreateStorageInterval = time.Hour t.Cleanup(func() { retryCreateStorageInterval = oldInterval }) createAttempted := make(chan struct{}, 1) + var createAttempts atomic.Int32 runTNStoreTest(t, func(s *store) { s.options.logServiceClientFactory = func(metadata.TNShard) (logservice.Client, error) { + createAttempts.Add(1) select { case createAttempted <- struct{}{}: default: @@ -273,13 +343,21 @@ func TestReplicaCreateRetryStopsPromptly(t *testing.T) { require.NoError(t, s.StartTNReplica(newTestTNShard(102, 202, 302))) select { case <-createAttempted: - case <-time.After(time.Second): + case <-time.After(storeTestTimeout): t.Fatal("storage creation was not attempted") } - start := time.Now() - s.stopper.Stop() - require.Less(t, time.Since(start), 500*time.Millisecond) + stopped := make(chan struct{}) + go func() { + s.stopper.Stop() + close(stopped) + }() + select { + case <-stopped: + case <-time.After(storeTestTimeout): + t.Fatal("store stop did not cancel replica creation retry") + } + require.Equal(t, int32(1), createAttempts.Load()) }) } diff --git a/pkg/txn/service/service.go b/pkg/txn/service/service.go index d64ca9dd853a6..8565ea203d3f5 100644 --- a/pkg/txn/service/service.go +++ b/pkg/txn/service/service.go @@ -17,7 +17,6 @@ package service import ( "bytes" "context" - "errors" "fmt" "sync" "time" @@ -39,10 +38,11 @@ import ( var _ TxnService = (*service)(nil) type service struct { - sid string - logger *log.MOLogger - shard metadata.TNShard - storage storage.TxnStorage + sid string + logger *log.MOLogger + shard metadata.TNShard + storage storage.TxnStorage + // sender is owned by the TN store and shared by all replica services. sender rpc.TxnSender stopper *stopper.Stopper allocator lockservice.LockTableAllocator @@ -127,7 +127,7 @@ func (s *service) Close(destroy bool) error { closer = s.storage.Destroy } // FIXME: all context.TODO() need to use tracing context - return errors.Join(closer(context.TODO()), s.sender.Close()) + return closer(context.TODO()) } func (s *service) gcZombieTxn(ctx context.Context) { @@ -255,7 +255,7 @@ func (s *service) parallelSendWithRetry( if err != nil { err = moerr.AttachCause(ctx, err) util.LogTxnSendRequestsFailed(s.logger, requests, err) - if !waitParallelSendRetryBackoff(ctx, backoff) { + if !waitRetryBackoff(ctx, backoff) { return nil } backoff = nextParallelSendRetryBackoff(backoff, maxBackoff) @@ -275,7 +275,7 @@ func (s *service) parallelSendWithRetry( return result } result.Release() - if !waitParallelSendRetryBackoff(ctx, backoff) { + if !waitRetryBackoff(ctx, backoff) { return nil } backoff = nextParallelSendRetryBackoff(backoff, maxBackoff) @@ -283,7 +283,7 @@ func (s *service) parallelSendWithRetry( } } -func waitParallelSendRetryBackoff(ctx context.Context, backoff time.Duration) bool { +func waitRetryBackoff(ctx context.Context, backoff time.Duration) bool { if backoff <= 0 { return ctx.Err() == nil } diff --git a/pkg/txn/service/service_cn_handler.go b/pkg/txn/service/service_cn_handler.go index 5b6ebeb3ebaf3..c216bc2621fdf 100644 --- a/pkg/txn/service/service_cn_handler.go +++ b/pkg/txn/service/service_cn_handler.go @@ -53,7 +53,9 @@ func (s *service) Read(ctx context.Context, request *txn.TxnRequest, response *t return nil } - s.waitClockTo(request.Txn.SnapshotTS) + if err := s.waitClockTo(ctx, request.Txn.SnapshotTS); err != nil { + return err + } // We do not write transaction information to sync.Map during read operations because commit and abort // for read-only transactions are not sent to the TN node, so there is no way to clean up the transaction @@ -445,7 +447,9 @@ func (s *service) startAsyncCommitTask(txnCtx *txnContext) error { } util.LogTxnCommittingFailed(s.logger, txnMeta, err) // TODO: make config - time.Sleep(time.Second) + if !waitRetryBackoff(ctx, time.Second) { + return + } } } @@ -495,12 +499,22 @@ func (s *service) checkCNRequest(request *txn.TxnRequest) { } } -func (s *service) waitClockTo(ts timestamp.Timestamp) { +func (s *service) waitClockTo(ctx context.Context, ts timestamp.Timestamp) error { for { now, _ := runtime.ServiceRuntime(s.sid).Clock().Now() if now.GreaterEq(ts) { - return + return nil + } + wait := time.Duration(ts.PhysicalTime - now.PhysicalTime) + if wait < time.Duration(math.MaxInt64) { + wait++ + } + timer := time.NewTimer(wait) + select { + case <-ctx.Done(): + timer.Stop() + return context.Cause(ctx) + case <-timer.C: } - time.Sleep(time.Duration(ts.PhysicalTime + 1 - now.PhysicalTime)) } } diff --git a/pkg/txn/service/service_cn_handler_test.go b/pkg/txn/service/service_cn_handler_test.go index 7fa73c8a57538..9a467b0509fcf 100644 --- a/pkg/txn/service/service_cn_handler_test.go +++ b/pkg/txn/service/service_cn_handler_test.go @@ -16,8 +16,10 @@ package service import ( "context" + "errors" "math" "os" + "sync" "sync/atomic" "testing" "time" @@ -33,6 +35,7 @@ import ( "github.com/matrixorigin/matrixone/pkg/pb/timestamp" "github.com/matrixorigin/matrixone/pkg/pb/txn" "github.com/matrixorigin/matrixone/pkg/txn/rpc" + "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" @@ -40,6 +43,24 @@ import ( "go.uber.org/zap/zapcore" ) +type committingErrorTxnStorage struct { + storage.TxnStorage + entered chan struct{} + calls atomic.Int32 +} + +func (s *committingErrorTxnStorage) Committing(ctx context.Context, _ txn.TxnMeta) error { + s.calls.Add(1) + select { + case s.entered <- struct{}{}: + default: + } + <-ctx.Done() + return errors.New("committing failed") +} + +const cancellationTestTimeout = 30 * time.Second + func TestReadBasic(t *testing.T) { sender := NewTestSender() defer func() { @@ -129,6 +150,84 @@ func TestReadBlockWithClock(t *testing.T) { assert.Equal(t, int64(3), ts) } +func TestWaitClockToHonorsContext(t *testing.T) { + sender := NewTestSender() + t.Cleanup(func() { require.NoError(t, sender.Close()) }) + clockEntered := make(chan struct{}) + releaseClock := make(chan struct{}) + var observeClock atomic.Bool + var signalClock sync.Once + var releaseClockOnce sync.Once + releaseClockWait := func() { + releaseClockOnce.Do(func() { close(releaseClock) }) + } + defer releaseClockWait() + clock := NewTestSpecClock(func() int64 { + if observeClock.Load() { + signalClock.Do(func() { close(clockEntered) }) + <-releaseClock + } + return 1 + }) + s := NewTestTxnService(t, 1, sender, clock).(*service) + require.NoError(t, s.Start()) + t.Cleanup(func() { require.NoError(t, s.Close(false)) }) + + cause := errors.New("clock wait canceled") + ctx, cancel := context.WithCancelCause(context.Background()) + defer cancel(cause) + done := make(chan error, 1) + observeClock.Store(true) + go func() { + done <- s.waitClockTo(ctx, NewTestTimestamp(math.MaxInt64)) + }() + select { + case <-clockEntered: + case <-time.After(cancellationTestTimeout): + t.Fatal("clock wait did not read the current timestamp") + } + cancel(cause) + releaseClockWait() + select { + case err := <-done: + require.ErrorIs(t, err, cause) + case <-time.After(cancellationTestTimeout): + t.Fatal("clock wait did not stop after context cancellation") + } +} + +func TestAsyncCommitRetryStopsOnClose(t *testing.T) { + sender := NewTestSender() + t.Cleanup(func() { require.NoError(t, sender.Close()) }) + s := NewTestTxnService(t, 1, sender, NewTestClock(0)).(*service) + storage := &committingErrorTxnStorage{ + TxnStorage: s.storage, + entered: make(chan struct{}, 1), + } + s.storage = storage + require.NoError(t, s.Start()) + + txnCtx, _ := s.maybeAddTxn(NewTestTxn(1, 1, 1)) + require.NoError(t, s.startAsyncCommitTask(txnCtx)) + select { + case <-storage.entered: + case <-time.After(cancellationTestTimeout): + t.Fatal("async commit did not call storage.Committing") + } + + closed := make(chan error, 1) + go func() { + closed <- s.Close(false) + }() + select { + case err := <-closed: + require.NoError(t, err) + case <-time.After(cancellationTestTimeout): + t.Fatal("service close was blocked by committing retry") + } + require.Equal(t, int32(1), storage.calls.Load()) +} + func TestReadCannotBlockByUncommitted(t *testing.T) { sender := NewTestSender() defer func() { diff --git a/pkg/txn/service/service_test.go b/pkg/txn/service/service_test.go index 404799d4d457a..fb2c9d4525c5d 100644 --- a/pkg/txn/service/service_test.go +++ b/pkg/txn/service/service_test.go @@ -16,6 +16,7 @@ package service import ( "context" + "sync/atomic" "testing" "time" @@ -30,6 +31,16 @@ type retryTestSender struct { send func(context.Context, []txn.TxnRequest) (*rpc.SendResult, error) } +type closeTrackingSender struct { + retryTestSender + closed atomic.Int32 +} + +func (s *closeTrackingSender) Close() error { + s.closed.Add(1) + return nil +} + func (s *retryTestSender) Send(ctx context.Context, requests []txn.TxnRequest) (*rpc.SendResult, error) { return s.send(ctx, requests) } @@ -38,6 +49,17 @@ func (s *retryTestSender) Close() error { return nil } +func TestTxnServiceDoesNotCloseBorrowedSender(t *testing.T) { + sender := &closeTrackingSender{} + sender.send = func(context.Context, []txn.TxnRequest) (*rpc.SendResult, error) { + return &rpc.SendResult{}, nil + } + s := NewTestTxnService(t, 1, sender, NewTestClock(1)) + assert.NoError(t, s.Start()) + assert.NoError(t, s.Close(false)) + assert.Zero(t, sender.closed.Load()) +} + func TestGCZombie(t *testing.T) { sender := NewTestSender() defer func() { diff --git a/pkg/txn/storage/mem/kv_txn_storage.go b/pkg/txn/storage/mem/kv_txn_storage.go index 0d621644a21b7..ba6a79d6b34e1 100644 --- a/pkg/txn/storage/mem/kv_txn_storage.go +++ b/pkg/txn/storage/mem/kv_txn_storage.go @@ -89,6 +89,8 @@ type Event struct { type KVTxnStorage struct { sync.RWMutex logClient logservice.Client + closeOnce sync.Once + closeErr error clock clock.Clock recoverFrom logservice.Lsn uncommittedTxn map[string]*txn.TxnMeta @@ -98,7 +100,7 @@ type KVTxnStorage struct { eventC chan Event } -// NewKVTxnStorage create KV-based implementation of TxnStorage +// NewKVTxnStorage creates KV-based TxnStorage and takes ownership of logClient. func NewKVTxnStorage(recoverFrom logservice.Lsn, logClient logservice.Client, clock clock.Clock) *KVTxnStorage { return &KVTxnStorage{ logClient: logClient, @@ -190,11 +192,14 @@ func (kv *KVTxnStorage) Start() error { } func (kv *KVTxnStorage) Close(ctx context.Context) error { - return nil + kv.closeOnce.Do(func() { + kv.closeErr = kv.logClient.Close() + }) + return kv.closeErr } func (kv *KVTxnStorage) Destroy(ctx context.Context) error { - return nil + return kv.Close(ctx) } func (kv *KVTxnStorage) Read(ctx context.Context, txnMeta txn.TxnMeta, op uint32, payload []byte) (storage.ReadResult, error) { diff --git a/pkg/txn/storage/mem/kv_txn_storage_test.go b/pkg/txn/storage/mem/kv_txn_storage_test.go index 198c61d3f40a8..f5c5d686139c0 100644 --- a/pkg/txn/storage/mem/kv_txn_storage_test.go +++ b/pkg/txn/storage/mem/kv_txn_storage_test.go @@ -30,6 +30,25 @@ import ( "github.com/stretchr/testify/assert" ) +type closeTrackingLogClient struct { + logservice.Client + closed atomic.Int32 +} + +func (c *closeTrackingLogClient) Close() error { + c.closed.Add(1) + return c.Client.Close() +} + +func TestCloseClosesLogClientOnce(t *testing.T) { + client := &closeTrackingLogClient{Client: NewMemLog()} + storage := NewKVTxnStorage(0, client, newTestClock(1)) + + assert.NoError(t, storage.Close(context.Background())) + assert.NoError(t, storage.Destroy(context.Background())) + assert.Equal(t, int32(1), client.closed.Load()) +} + func TestWrite(t *testing.T) { l := NewMemLog() s := NewKVTxnStorage(0, l, newTestClock(1)) diff --git a/pkg/txn/storage/tae/storage.go b/pkg/txn/storage/tae/storage.go index b23012856bd06..b584f1a05f65a 100644 --- a/pkg/txn/storage/tae/storage.go +++ b/pkg/txn/storage/tae/storage.go @@ -109,6 +109,9 @@ func newTAEStorage( if rt.ServiceUUID() != opt.SID { panic(fmt.Sprintf("service uuid mismatch, %s != %s", rt.ServiceUUID(), opt.SID)) } + if err := service.ValidateRPCMaxMessageSize(logtailServerCfg.RpcMaxMessageSize); err != nil { + return nil, err + } taeHandler, err := deps.newTAEHandle(ctx, dataDir, client, opt) if err != nil { return nil, err diff --git a/pkg/txn/storage/tae/storage_lifecycle_test.go b/pkg/txn/storage/tae/storage_lifecycle_test.go index 0bda2c5b1beaa..54d4877612422 100644 --- a/pkg/txn/storage/tae/storage_lifecycle_test.go +++ b/pkg/txn/storage/tae/storage_lifecycle_test.go @@ -111,6 +111,32 @@ func TestNewTAEStorageHandleCreationFailure(t *testing.T) { require.Equal(t, 0, serverCalls) } +func TestNewTAEStorageRejectsInvalidLogtailMessageSizeBeforeOpeningHandle(t *testing.T) { + handleCalls := 0 + deps := taeStorageDependencies{ + newTAEHandle: func( + context.Context, + string, + client.QueryClient, + *options.Options, + ) (taeHandle, error) { + handleCalls++ + return nil, nil + }, + } + rt := runtime.DefaultRuntime() + cfg := options.NewDefaultLogtailServerCfg() + cfg.RpcMaxMessageSize = 1 + + storage, err := newTAEStorage( + context.Background(), t.TempDir(), + &options.Options{SID: rt.ServiceUUID()}, metadata.TNShard{}, rt, + "", cfg, nil, nil, deps) + require.Nil(t, storage) + require.Error(t, err) + require.Zero(t, handleCalls) +} + func TestNewTAEStorageLogtailServerFailureClosesHandle(t *testing.T) { primaryErr := errors.New("logtail server creation failed") cleanupErr := errors.New("handle close failed") @@ -158,7 +184,7 @@ func newTAEStorageForTest( metadata.TNShard{}, rt, "", - &options.LogtailServerCfg{}, + options.NewDefaultLogtailServerCfg(), nil, nil, deps, diff --git a/pkg/vm/engine/tae/logtail/service/response.go b/pkg/vm/engine/tae/logtail/service/response.go index 9c295c42dddbd..9cfde0e5d7df2 100644 --- a/pkg/vm/engine/tae/logtail/service/response.go +++ b/pkg/vm/engine/tae/logtail/service/response.go @@ -17,12 +17,47 @@ package service import ( "fmt" "math" + "math/bits" "sync" + "github.com/matrixorigin/matrixone/pkg/common/moerr" "github.com/matrixorigin/matrixone/pkg/common/morpc" "github.com/matrixorigin/matrixone/pkg/pb/logtail" ) +const maxRPCMessageSize = 100 * 1024 * 1024 + +// ValidateRPCMaxMessageSize verifies that response segments have room for a +// payload and cannot request an unsafe allocation. +func ValidateRPCMaxMessageSize(maxMessageSize int64) error { + if maxMessageSize <= 0 { + return moerr.NewBadConfigNoCtxf( + "logtail rpc max message size must be positive, got %d", maxMessageSize) + } + if maxMessageSize > maxRPCMessageSize { + return moerr.NewBadConfigNoCtxf( + "logtail rpc max message size %d exceeds the supported limit %d", + maxMessageSize, maxRPCMessageSize) + } + if leastEffectiveCapacity(int(maxMessageSize)) <= 0 { + return moerr.NewBadConfigNoCtxf( + "logtail rpc max message size %d is too small", maxMessageSize) + } + return nil +} + +func leastEffectiveCapacity(maxMessageSize int) int { + segment := LogtailResponseSegment{} + segment.StreamID = math.MaxUint64 + segment.Sequence = math.MaxInt32 + segment.MaxSequence = math.MaxInt32 + segment.MessageSize = math.MaxInt32 + // Payload contributes one field tag and the encoded length. Compute the + // worst-case header without allocating a maxMessageSize payload. + headerSize := segment.ProtoSize() + 1 + (bits.Len64(uint64(maxMessageSize)|1)+6)/7 + return maxMessageSize - headerSize +} + // LogtailPhase is the logtail information of one phase of // subscription request. Phase 1 is executed asynchronously // to collect most of logtails of the subscription request. @@ -165,15 +200,5 @@ func (p *serverSegmentPool) Release(seg *LogtailResponseSegment) { } func (p *serverSegmentPool) LeastEffectiveCapacity() int { - segment := p.Acquire() - defer p.Release(segment) - - segment.StreamID = math.MaxUint64 - segment.Sequence = math.MaxInt32 - segment.MaxSequence = math.MaxInt32 - segment.MessageSize = math.MaxInt32 - maxHeaderSize := segment.ProtoSize() - p.maxMessageSize - - // Take out reserved size, then effective capacity left. - return p.maxMessageSize - maxHeaderSize + return leastEffectiveCapacity(p.maxMessageSize) } diff --git a/pkg/vm/engine/tae/logtail/service/response_test.go b/pkg/vm/engine/tae/logtail/service/response_test.go index 9265ccabcc89f..45119467b7bc9 100644 --- a/pkg/vm/engine/tae/logtail/service/response_test.go +++ b/pkg/vm/engine/tae/logtail/service/response_test.go @@ -21,6 +21,13 @@ import ( "github.com/stretchr/testify/require" ) +func TestValidateRPCMaxMessageSize(t *testing.T) { + require.Error(t, ValidateRPCMaxMessageSize(1)) + require.Error(t, ValidateRPCMaxMessageSize(math.MaxInt64)) + require.NoError(t, ValidateRPCMaxMessageSize(maxRPCMessageSize)) + require.Error(t, ValidateRPCMaxMessageSize(maxRPCMessageSize+1)) +} + func TestResponseSize(t *testing.T) { maxMessageSize := 1024 pool := NewLogtailServerSegmentPool(maxMessageSize) diff --git a/pkg/vm/engine/tae/logtail/service/server.go b/pkg/vm/engine/tae/logtail/service/server.go index c1aee95704d26..1faeb0d4819dc 100644 --- a/pkg/vm/engine/tae/logtail/service/server.go +++ b/pkg/vm/engine/tae/logtail/service/server.go @@ -190,6 +190,9 @@ func NewLogtailServer( for _, opt := range opts { opt(s) } + if err := ValidateRPCMaxMessageSize(s.cfg.RpcMaxMessageSize); err != nil { + return nil, err + } uid, _ := uuid.NewV7() s.logger = s.logger.Named(LogtailServiceRPCName). @@ -200,7 +203,8 @@ func NewLogtailServer( s.pool.segments = NewLogtailServerSegmentPool(int(s.cfg.RpcMaxMessageSize)) s.maxChunkSize = s.pool.segments.LeastEffectiveCapacity() if s.maxChunkSize <= 0 { - panic("rpc max message size isn't enough") + return nil, moerr.NewBadConfigNoCtxf( + "logtail rpc max message size %d is too small", s.cfg.RpcMaxMessageSize) } s.logger.Debug("max data chunk size for segment", zap.Int("value", s.maxChunkSize)) diff --git a/pkg/vm/engine/tae/logtail/service/server_test.go b/pkg/vm/engine/tae/logtail/service/server_test.go index 411ca905baecb..c42a216c98e68 100644 --- a/pkg/vm/engine/tae/logtail/service/server_test.go +++ b/pkg/vm/engine/tae/logtail/service/server_test.go @@ -122,6 +122,30 @@ func TestService(t *testing.T) { } } +func TestNewLogtailServerRejectsSmallRPCMessageSize(t *testing.T) { + cfg := options.NewDefaultLogtailServerCfg() + cfg.RpcMaxMessageSize = 1 + + require.NotPanics(t, func() { + server, err := NewLogtailServer( + "127.0.0.1:0", cfg, mockLocktailer(), mockRuntime(), nil) + require.Nil(t, server) + require.Error(t, err) + }) +} + +func TestNewLogtailServerValidatesFinalOptionMessageSize(t *testing.T) { + cfg := options.NewDefaultLogtailServerCfg() + + require.NotPanics(t, func() { + server, err := NewLogtailServer( + "127.0.0.1:0", cfg, mockLocktailer(), mockRuntime(), nil, + WithServerMaxMessageSize(1)) + require.Nil(t, server) + require.Error(t, err) + }) +} + type logtailer struct { tables []api.TableID }