This is an automated email from the ASF dual-hosted git repository.

AlexStocks pushed a commit to branch develop
in repository https://gitbox.apache.org/repos/asf/dubbo-go.git


The following commit(s) were added to refs/heads/develop by this push:
     new 4649e8592 fix: make lifecycle waits cancellable (#3599)
4649e8592 is described below

commit 4649e859233ee226951fa25aa9b38161aa9a7e80
Author: DadaVinqi <[email protected]>
AuthorDate: Tue Aug 11 11:16:58 2026 +0800

    fix: make lifecycle waits cancellable (#3599)
    
    * fix: make lifecycle waits cancellable
    
    Signed-off-by: DadaVinqi <[email protected]>
    
    * fix: refine shutdown wait semantics
    
    Signed-off-by: DadaVinqi <[email protected]>
    
    * refactor: reduce shutdown code complexity
    
    Signed-off-by: DadaVinqi <[email protected]>
    
    * test: cover getty close before publish
    
    Signed-off-by: DadaVinqi <[email protected]>
    
    ---------
    
    Signed-off-by: DadaVinqi <[email protected]>
    Co-authored-by: DadaVinqi <[email protected]>
---
 cluster/cluster/failback/cluster_invoker.go | 242 ++++++++++++++++++++++-----
 cluster/cluster/failback/cluster_test.go    | 251 ++++++++++++++++++++++++++++
 registry/base_registry.go                   |  58 +++++--
 registry/base_registry_test.go              | 208 +++++++++++++++++++++++
 remoting/getty/getty_client.go              | 123 +++++++++-----
 remoting/getty/getty_client_test.go         | 144 ++++++++++++++++
 remoting/getty/pool.go                      |  49 +++---
 7 files changed, 948 insertions(+), 127 deletions(-)

diff --git a/cluster/cluster/failback/cluster_invoker.go 
b/cluster/cluster/failback/cluster_invoker.go
index 3509aad83..6cb18eac9 100644
--- a/cluster/cluster/failback/cluster_invoker.go
+++ b/cluster/cluster/failback/cluster_invoker.go
@@ -19,6 +19,7 @@ package failback
 
 import (
        "context"
+       "errors"
        "strconv"
        "sync"
        "time"
@@ -51,13 +52,21 @@ import (
 type failbackClusterInvoker struct {
        base.BaseClusterInvoker
 
-       once          sync.Once
-       ticker        *time.Ticker
        maxRetries    int64
        failbackTasks int64
+
+       lifecycleMu   sync.Mutex
+       stopped       bool
        taskList      *queue.Queue
+       retryCancel   context.CancelFunc
+       processDone   chan struct{}
+       retryDone     chan struct{}
+       activeRetries int
+       destroyOnce   sync.Once
 }
 
+var errFailbackInvokerStopped = errors.New("failback invoker is stopped")
+
 func newFailbackClusterInvoker(directory directory.Directory) 
protocolbase.Invoker {
        invoker := &failbackClusterInvoker{
                BaseClusterInvoker: base.NewBaseClusterInvoker(directory),
@@ -79,49 +88,88 @@ func newFailbackClusterInvoker(directory 
directory.Directory) protocolbase.Invok
 }
 
 func (invoker *failbackClusterInvoker) tryTimerTaskProc(ctx context.Context, 
retryTask *retryTimerTask) {
+       if ctx.Err() != nil {
+               return
+       }
+
        invoked := make([]protocolbase.Invoker, 0)
        invoked = append(invoked, retryTask.lastInvoker)
 
        retryInvoker := invoker.DoSelect(retryTask.loadbalance, 
retryTask.invocation, retryTask.invokers, invoked)
+       if retryInvoker == nil || ctx.Err() != nil {
+               return
+       }
+
        res := retryInvoker.Invoke(ctx, retryTask.invocation)
-       if res.Error() != nil {
+       if res.Error() != nil && ctx.Err() == nil {
                retryTask.lastInvoker = retryInvoker
                retryTask.lastErr = res.Error()
                retryTask.checkRetry()
        }
 }
 
-func (invoker *failbackClusterInvoker) process(ctx context.Context) {
-       invoker.ticker = time.NewTicker(time.Second * 1)
-       for range invoker.ticker.C {
-               // check each timeout task and re-run
-               for {
-                       value, err := invoker.taskList.Peek()
-                       if err == queue.ErrDisposed {
+func (invoker *failbackClusterInvoker) process(ctx context.Context, taskList 
*queue.Queue, done chan struct{}) {
+       defer close(done)
+
+       ticker := time.NewTicker(time.Second)
+       defer ticker.Stop()
+
+       for {
+               select {
+               case <-ctx.Done():
+                       return
+               case <-ticker.C:
+                       if invoker.processRetryTasks(ctx, taskList) {
                                return
                        }
-                       if err == queue.ErrEmptyQueue {
-                               break
-                       }
+               }
+       }
+}
 
-                       retryTask := value.(*retryTimerTask)
-                       // use exponential backoff calculated wait time instead 
of fixed 5 seconds
-                       if time.Since(retryTask.lastT) < retryTask.nextBackoff {
-                               break
-                       }
+func (invoker *failbackClusterInvoker) processRetryTasks(ctx context.Context, 
taskList *queue.Queue) bool {
+       for {
+               select {
+               case <-ctx.Done():
+                       return true
+               default:
+               }
 
-                       // ignore return. the get must success.
-                       if _, err = invoker.taskList.Get(1); err != nil {
-                               logger.Warnf("[Cluster][Failback] get task 
failed, err=%v", err)
-                               break
-                       }
-                       go invoker.tryTimerTaskProc(ctx, retryTask)
+               value, err := taskList.Peek()
+               if err == queue.ErrDisposed {
+                       return true
+               }
+               if err == queue.ErrEmptyQueue {
+                       return false
                }
+               if err != nil {
+                       logger.Warnf("[Cluster][Failback] peek task failed, 
err=%v", err)
+                       return false
+               }
+
+               retryTask := value.(*retryTimerTask)
+               // use exponential backoff calculated wait time instead of 
fixed 5 seconds
+               if time.Since(retryTask.lastT) < retryTask.nextBackoff {
+                       return false
+               }
+
+               // ignore return. the get must success.
+               if _, err = taskList.Get(1); err != nil {
+                       logger.Warnf("[Cluster][Failback] get task failed, 
err=%v", err)
+                       return false
+               }
+               invoker.startRetry(ctx, retryTask)
        }
 }
 
 // Invoke executes with failback semantics: schedule retries on failure.
 func (invoker *failbackClusterInvoker) Invoke(ctx context.Context, invocation 
protocolbase.Invocation) result.Result {
+       if invoker.isStopped() {
+               return &result.RPCResult{Err: errFailbackInvokerStopped}
+       }
+       if err := invoker.CheckWhetherDestroyed(); err != nil {
+               return &result.RPCResult{Err: err}
+       }
+
        invokers := invoker.Directory.List(invocation)
        if err := invoker.CheckInvokers(invokers, invocation); err != nil {
                logger.Errorf("[Cluster][Failback] check invokers failed, 
method=%s service=%s err=%v",
@@ -144,19 +192,8 @@ func (invoker *failbackClusterInvoker) Invoke(ctx 
context.Context, invocation pr
        // DO INVOKE
        res := ivk.Invoke(ctx, invocation)
        if res.Error() != nil {
-               invoker.once.Do(func() {
-                       invoker.taskList = queue.New(invoker.failbackTasks)
-                       go invoker.process(ctx)
-               })
-
-               taskLen := invoker.taskList.Len()
-               if taskLen >= invoker.failbackTasks {
-                       logger.Warnf("[Cluster][Failback] task list full, 
len=%d", taskLen)
-                       return &result.RPCResult{}
-               }
-
                timerTask := newRetryTimerTask(loadBalance, invocation, 
invokers, ivk, invoker)
-               invoker.taskList.Put(timerTask)
+               invoker.enqueueInitialRetry(ctx, timerTask)
 
                logger.Errorf("[Cluster][Failback] invoke failed, method=%s 
service=%s err=%v",
                        methodName, url.Service(), res.Error().Error())
@@ -166,15 +203,131 @@ func (invoker *failbackClusterInvoker) Invoke(ctx 
context.Context, invocation pr
        return res
 }
 
+func (invoker *failbackClusterInvoker) isStopped() bool {
+       invoker.lifecycleMu.Lock()
+       defer invoker.lifecycleMu.Unlock()
+       return invoker.stopped
+}
+
 func (invoker *failbackClusterInvoker) Destroy() {
-       invoker.BaseClusterInvoker.Destroy()
+       invoker.destroyOnce.Do(func() {
+               invoker.lifecycleMu.Lock()
+               invoker.stopped = true
+               if invoker.retryCancel != nil {
+                       invoker.retryCancel()
+               }
+               taskList := invoker.taskList
+               processDone := invoker.processDone
+               retryDone := invoker.retryDone
+               if taskList != nil {
+                       _ = taskList.Dispose()
+               }
+               invoker.lifecycleMu.Unlock()
+
+               invoker.waitForShutdown(processDone, retryDone)
+               invoker.BaseClusterInvoker.Destroy()
+       })
+}
+
+func (invoker *failbackClusterInvoker) enqueueInitialRetry(ctx 
context.Context, retryTask *retryTimerTask) {
+       invoker.lifecycleMu.Lock()
+       defer invoker.lifecycleMu.Unlock()
+
+       if invoker.stopped || invoker.Destroyed.Load() {
+               return
+       }
+
+       if invoker.taskList == nil {
+               if ctx == nil {
+                       ctx = context.Background()
+               }
+               retryCtx, retryCancel := 
context.WithCancel(context.WithoutCancel(ctx))
+               invoker.retryCancel = retryCancel
+               invoker.taskList = queue.New(invoker.failbackTasks)
+               invoker.processDone = make(chan struct{})
+               go invoker.process(retryCtx, invoker.taskList, 
invoker.processDone)
+       }
+
+       if invoker.taskList.Len() >= invoker.failbackTasks {
+               logger.Warnf("[Cluster][Failback] task list full, len=%d", 
invoker.taskList.Len())
+               return
+       }
+
+       if err := invoker.taskList.Put(retryTask); err != nil {
+               logger.Warnf("[Cluster][Failback] put initial task failed, 
err=%v", err)
+       }
+}
+
+func (invoker *failbackClusterInvoker) startRetry(ctx context.Context, 
retryTask *retryTimerTask) {
+       invoker.lifecycleMu.Lock()
+       defer invoker.lifecycleMu.Unlock()
+
+       if invoker.stopped || ctx.Err() != nil {
+               return
+       }
+       if invoker.activeRetries == 0 {
+               invoker.retryDone = make(chan struct{})
+       }
+       invoker.activeRetries++
+       retryDone := invoker.retryDone
+       go func() {
+               defer invoker.finishRetry(retryDone)
+               invoker.tryTimerTaskProc(ctx, retryTask)
+       }()
+}
+
+func (invoker *failbackClusterInvoker) finishRetry(retryDone chan struct{}) {
+       invoker.lifecycleMu.Lock()
+       defer invoker.lifecycleMu.Unlock()
+
+       invoker.activeRetries--
+       if invoker.activeRetries == 0 && invoker.retryDone == retryDone {
+               close(retryDone)
+       }
+}
+
+func (invoker *failbackClusterInvoker) enqueueRetry(retryTask *retryTimerTask) 
bool {
+       invoker.lifecycleMu.Lock()
+       defer invoker.lifecycleMu.Unlock()
 
-       // stop ticker
-       if invoker.ticker != nil {
-               invoker.ticker.Stop()
+       if invoker.stopped || invoker.taskList == nil {
+               return false
        }
 
-       _ = invoker.taskList.Dispose()
+       retryTask.lastT = time.Now()
+       if err := invoker.taskList.Put(retryTask); err != nil {
+               logger.Warnf("[Cluster][Failback] put retry task failed, 
err=%v", err)
+               return false
+       }
+       return true
+}
+
+func (invoker *failbackClusterInvoker) waitForShutdown(processDone, retryDone 
<-chan struct{}) {
+       if processDone == nil && retryDone == nil {
+               return
+       }
+
+       wait := func(done <-chan struct{}, name string) bool {
+               if done == nil {
+                       return true
+               }
+
+               timer := 
time.NewTimer(constant.DefaultShutdownConfigStepTimeout)
+               defer timer.Stop()
+
+               select {
+               case <-done:
+                       return true
+               case <-timer.C:
+                       logger.Warnf("[Cluster][Failback] timed out waiting for 
%s shutdown", name)
+                       return false
+               }
+       }
+
+       if !wait(processDone, "retry processor") {
+               return
+       }
+       _ = wait(retryDone, "retry tasks")
 }
 
 type retryTimerTask struct {
@@ -203,13 +356,10 @@ func (t *retryTimerTask) checkRetry() {
                return
        }
 
-       logger.Infof("[Cluster][Failback] retry scheduled, backoff=%v 
method=%s", t.nextBackoff, t.invocation.MethodName())
-
-       if err := t.clusterInvoker.taskList.Put(t); err != nil {
-               logger.Errorf("[Cluster][Failback] put task failed, task=%v 
err=%v", t, err)
+       if !t.clusterInvoker.enqueueRetry(t) {
                return
        }
-       t.lastT = time.Now() // update lastT after successful Put
+       logger.Infof("[Cluster][Failback] retry scheduled, backoff=%v 
method=%s", t.nextBackoff, t.invocation.MethodName())
 }
 
 func newRetryTimerTask(loadbalance loadbalance.LoadBalance, invocation 
protocolbase.Invocation, invokers []protocolbase.Invoker,
diff --git a/cluster/cluster/failback/cluster_test.go 
b/cluster/cluster/failback/cluster_test.go
index c7a019354..f015328ce 100644
--- a/cluster/cluster/failback/cluster_test.go
+++ b/cluster/cluster/failback/cluster_test.go
@@ -84,6 +84,257 @@ func TestFailbackSuceess(t *testing.T) {
 
        result := clusterInvoker.Invoke(context.Background(), 
&invocation.RPCInvocation{})
        assert.Equal(t, mockResult, result)
+
+       invoker.EXPECT().Destroy().Return()
+       clusterInvoker.Destroy()
+       clusterInvoker.Destroy()
+}
+
+func TestFailbackDestroyWithoutFailure(t *testing.T) {
+       ctrl := gomock.NewController(t)
+       defer ctrl.Finish()
+
+       invoker := mock.NewMockInvoker(ctrl)
+       clusterInvoker := registerFailback(invoker).(*failbackClusterInvoker)
+
+       invoker.EXPECT().Destroy().Return()
+       require.NotPanics(t, clusterInvoker.Destroy)
+       require.NotPanics(t, clusterInvoker.Destroy)
+}
+
+func TestFailbackInvokeAfterDestroy(t *testing.T) {
+       ctrl := gomock.NewController(t)
+       defer ctrl.Finish()
+
+       invoker := mock.NewMockInvoker(ctrl)
+       clusterInvoker := registerFailback(invoker).(*failbackClusterInvoker)
+
+       invoker.EXPECT().Destroy().Return()
+       clusterInvoker.Destroy()
+
+       result := clusterInvoker.Invoke(context.Background(), 
&invocation.RPCInvocation{})
+       require.ErrorIs(t, result.Error(), errFailbackInvokerStopped)
+}
+
+func TestFailbackRetryUsesIndependentContext(t *testing.T) {
+       ctrl := gomock.NewController(t)
+       defer ctrl.Finish()
+
+       invoker := mock.NewMockInvoker(ctrl)
+       clusterInvoker := registerFailback(invoker).(*failbackClusterInvoker)
+
+       invoker.EXPECT().GetURL().Return(failbackUrl).AnyTimes()
+       invoker.EXPECT().IsAvailable().Return(true).AnyTimes()
+
+       failedResult := &result.RPCResult{Err: perrors.New("error")}
+       successResult := &result.RPCResult{Rest: clusterpkg.Rest{Tried: 0, 
Success: true}}
+       retryStarted := make(chan struct{})
+       retryContextErr := make(chan error, 1)
+       var callCount atomic.Int32
+
+       invoker.EXPECT().Invoke(gomock.Any(), 
gomock.Any()).Times(2).DoAndReturn(
+               func(ctx context.Context, _ base.Invocation) result.Result {
+                       if callCount.Add(1) == 1 {
+                               return failedResult
+                       }
+                       retryContextErr <- ctx.Err()
+                       close(retryStarted)
+                       return successResult
+               },
+       )
+
+       ctx, cancel := context.WithTimeout(context.Background(), 
50*time.Millisecond)
+       defer cancel()
+
+       result := clusterInvoker.Invoke(ctx, &invocation.RPCInvocation{})
+       require.NoError(t, result.Error())
+       <-ctx.Done()
+
+       select {
+       case <-retryStarted:
+       case <-time.After(5 * time.Second):
+               t.Fatal("failback retry did not start after caller context 
cancellation")
+       }
+       require.NoError(t, <-retryContextErr)
+
+       invoker.EXPECT().Destroy().Return()
+       clusterInvoker.Destroy()
+}
+
+func TestFailbackDestroyCancelsRetry(t *testing.T) {
+       ctrl := gomock.NewController(t)
+       defer ctrl.Finish()
+
+       invoker := mock.NewMockInvoker(ctrl)
+       clusterInvoker := registerFailback(invoker).(*failbackClusterInvoker)
+
+       invoker.EXPECT().GetURL().Return(failbackUrl).AnyTimes()
+       invoker.EXPECT().IsAvailable().Return(true).AnyTimes()
+
+       failedResult := &result.RPCResult{Err: perrors.New("error")}
+       retryStarted := make(chan struct{})
+       retryReturned := make(chan struct{})
+       var callCount atomic.Int32
+
+       invoker.EXPECT().Invoke(gomock.Any(), 
gomock.Any()).Times(2).DoAndReturn(
+               func(ctx context.Context, _ base.Invocation) result.Result {
+                       if callCount.Add(1) == 1 {
+                               return failedResult
+                       }
+                       close(retryStarted)
+                       <-ctx.Done()
+                       close(retryReturned)
+                       return &result.RPCResult{Err: ctx.Err()}
+               },
+       )
+
+       result := clusterInvoker.Invoke(context.Background(), 
&invocation.RPCInvocation{})
+       require.NoError(t, result.Error())
+
+       select {
+       case <-retryStarted:
+       case <-time.After(5 * time.Second):
+               t.Fatal("failback retry did not start")
+       }
+
+       invoker.EXPECT().Destroy().Return()
+       destroyed := make(chan struct{})
+       go func() {
+               clusterInvoker.Destroy()
+               close(destroyed)
+       }()
+
+       select {
+       case <-retryReturned:
+       case <-time.After(2 * time.Second):
+               t.Fatal("retry invocation did not observe shutdown 
cancellation")
+       }
+       select {
+       case <-destroyed:
+       case <-time.After(2 * time.Second):
+               t.Fatal("Destroy did not return after retry cancellation")
+       }
+
+       clusterInvoker.Destroy()
+}
+
+func TestFailbackDestroyHasBoundedRetryWait(t *testing.T) {
+       ctrl := gomock.NewController(t)
+       defer ctrl.Finish()
+
+       invoker := mock.NewMockInvoker(ctrl)
+       clusterInvoker := registerFailback(invoker).(*failbackClusterInvoker)
+
+       invoker.EXPECT().GetURL().Return(failbackUrl).AnyTimes()
+       invoker.EXPECT().IsAvailable().Return(true).AnyTimes()
+
+       failedResult := &result.RPCResult{Err: perrors.New("error")}
+       retryStarted := make(chan struct{})
+       retryReturned := make(chan struct{})
+       releaseRetry := make(chan struct{})
+       var releaseOnce sync.Once
+       release := func() {
+               releaseOnce.Do(func() {
+                       close(releaseRetry)
+               })
+       }
+       defer release()
+       var callCount atomic.Int32
+
+       invoker.EXPECT().Invoke(gomock.Any(), 
gomock.Any()).Times(2).DoAndReturn(
+               func(_ context.Context, _ base.Invocation) result.Result {
+                       if callCount.Add(1) == 1 {
+                               return failedResult
+                       }
+                       close(retryStarted)
+                       <-releaseRetry
+                       close(retryReturned)
+                       return &result.RPCResult{Rest: clusterpkg.Rest{Tried: 
0, Success: true}}
+               },
+       )
+
+       result := clusterInvoker.Invoke(context.Background(), 
&invocation.RPCInvocation{})
+       require.NoError(t, result.Error())
+
+       select {
+       case <-retryStarted:
+       case <-time.After(5 * time.Second):
+               t.Fatal("failback retry did not start")
+       }
+
+       invoker.EXPECT().Destroy().Return()
+       destroyed := make(chan struct{})
+       go func() {
+               clusterInvoker.Destroy()
+               close(destroyed)
+       }()
+
+       select {
+       case <-destroyed:
+       case <-time.After(5 * time.Second):
+               t.Fatal("Destroy blocked past the bounded retry wait")
+       }
+       release()
+
+       select {
+       case <-retryReturned:
+       case <-time.After(2 * time.Second):
+               t.Fatal("retry goroutine did not finish after release")
+       }
+}
+
+func TestFailbackWaitForShutdownUsesPerStepTimeout(t *testing.T) {
+       processDone := make(chan struct{})
+       retryDone := make(chan struct{})
+       waitDone := make(chan struct{})
+
+       go func() {
+               time.Sleep(constant.DefaultShutdownConfigStepTimeout - 
time.Second)
+               close(processDone)
+               time.Sleep(1500 * time.Millisecond)
+               close(retryDone)
+       }()
+       go func() {
+               (&failbackClusterInvoker{}).waitForShutdown(processDone, 
retryDone)
+               close(waitDone)
+       }()
+
+       select {
+       case <-waitDone:
+       case <-time.After(constant.DefaultShutdownConfigStepTimeout + 
time.Second):
+               t.Fatal("waitForShutdown did not return")
+       }
+
+       select {
+       case <-retryDone:
+       default:
+               t.Fatal("waitForShutdown returned before the retry tasks 
completed")
+       }
+}
+
+func TestFailbackDoesNotEnqueueAfterDestroy(t *testing.T) {
+       ctrl := gomock.NewController(t)
+       defer ctrl.Finish()
+
+       invoker := mock.NewMockInvoker(ctrl)
+       clusterInvoker := registerFailback(invoker).(*failbackClusterInvoker)
+
+       invoker.EXPECT().GetURL().Return(failbackUrl).AnyTimes()
+       invoker.EXPECT().IsAvailable().Return(true).AnyTimes()
+       invoker.EXPECT().Invoke(gomock.Any(), gomock.Any()).Return(
+               &result.RPCResult{Err: perrors.New("error")},
+       )
+
+       result := clusterInvoker.Invoke(context.Background(), 
&invocation.RPCInvocation{})
+       require.NoError(t, result.Error())
+
+       value, err := clusterInvoker.taskList.Peek()
+       require.NoError(t, err)
+       retryTask := value.(*retryTimerTask)
+
+       invoker.EXPECT().Destroy().Return()
+       clusterInvoker.Destroy()
+       require.False(t, clusterInvoker.enqueueRetry(retryTask))
 }
 
 // failed firstly, success later after one retry.
diff --git a/registry/base_registry.go b/registry/base_registry.go
index 29442066c..e12fdf04b 100644
--- a/registry/base_registry.go
+++ b/registry/base_registry.go
@@ -45,7 +45,8 @@ const (
 )
 
 var (
-       localIP = ""
+       localIP                    = ""
+       errBaseRegistryUnavailable = perrors.New("BaseRegistry is not 
available")
 )
 
 func init() {
@@ -330,30 +331,53 @@ func (r *BaseRegistry) Subscribe(url *common.URL, 
notifyListener NotifyListener)
        for {
                if !r.IsAvailable() {
                        logger.Warn("[Registry] event listener game over")
-                       return perrors.New("BaseRegistry is not available")
+                       return errBaseRegistryUnavailable
                }
 
                listener, err := r.facadeBasedRegistry.DoSubscribe(url)
+               if err == nil {
+                       return r.watchListener(listener, notifyListener)
+               }
+               if !r.IsAvailable() {
+                       logger.Warn("[Registry] event listener game over")
+                       return errBaseRegistryUnavailable
+               }
+               logger.Warnf("[Registry] getListener() = err=%v", 
perrors.WithStack(err))
+               if err = r.waitRetryDelay(); err != nil {
+                       return err
+               }
+       }
+}
+
+func (r *BaseRegistry) watchListener(listener Listener, notifyListener 
NotifyListener) error {
+       for {
+               serviceEvent, err := listener.Next()
                if err != nil {
+                       logger.Warnf("[Registry] Selector.watch() = err=%v", 
perrors.WithStack(err))
+                       listener.Close()
                        if !r.IsAvailable() {
-                               logger.Warn("[Registry] event listener game 
over")
-                               return err
+                               return errBaseRegistryUnavailable
                        }
-                       logger.Warnf("[Registry] getListener() = err=%v", 
perrors.WithStack(err))
-                       time.Sleep(time.Duration(RegistryConnDelay) * 
time.Second)
-                       continue
+                       return nil
                }
-
-               for {
-                       if serviceEvent, err := listener.Next(); err != nil {
-                               logger.Warnf("[Registry] Selector.watch() = 
err=%v", perrors.WithStack(err))
-                               listener.Close()
-                               return nil
-                       } else {
-                               logger.Debugf("[Registry] update begin, 
event=%v", serviceEvent.String())
-                               notifyListener.Notify(serviceEvent)
-                       }
+               if !r.IsAvailable() {
+                       listener.Close()
+                       return errBaseRegistryUnavailable
                }
+               logger.Debugf("[Registry] update begin, event=%v", 
serviceEvent.String())
+               notifyListener.Notify(serviceEvent)
+       }
+}
+
+func (r *BaseRegistry) waitRetryDelay() error {
+       timer := time.NewTimer(time.Duration(RegistryConnDelay) * time.Second)
+       defer timer.Stop()
+
+       select {
+       case <-r.Done():
+               return errBaseRegistryUnavailable
+       case <-timer.C:
+               return nil
        }
 }
 
diff --git a/registry/base_registry_test.go b/registry/base_registry_test.go
new file mode 100644
index 000000000..aac8c4bbd
--- /dev/null
+++ b/registry/base_registry_test.go
@@ -0,0 +1,208 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *     http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package registry
+
+import (
+       "errors"
+       "sync"
+       "sync/atomic"
+       "testing"
+       "time"
+)
+
+import (
+       "github.com/stretchr/testify/require"
+)
+
+import (
+       "dubbo.apache.org/dubbo-go/v3/common"
+)
+
+type baseRegistryTestFacade struct {
+       BaseRegistry
+
+       subscribeCalls   atomic.Int32
+       firstSubscribe   chan struct{}
+       subscribeRelease chan struct{}
+       listener         Listener
+       subscribeErr     error
+}
+
+func newBaseRegistryTestFacade(listener Listener, subscribeErr error) 
*baseRegistryTestFacade {
+       facade := &baseRegistryTestFacade{
+               firstSubscribe: make(chan struct{}),
+               listener:       listener,
+       }
+       facade.InitBaseRegistry(common.NewURLWithOptions(), facade)
+       facade.subscribeErr = subscribeErr
+       return facade
+}
+
+func (f *baseRegistryTestFacade) DoSubscribe(*common.URL) (Listener, error) {
+       if f.subscribeCalls.Add(1) == 1 {
+               close(f.firstSubscribe)
+       }
+       if f.subscribeRelease != nil {
+               <-f.subscribeRelease
+       }
+       return f.listener, f.subscribeErr
+}
+
+func (f *baseRegistryTestFacade) DoUnsubscribe(*common.URL) (Listener, error) {
+       return nil, nil
+}
+
+func (f *baseRegistryTestFacade) CreatePath(string) error {
+       return nil
+}
+
+func (f *baseRegistryTestFacade) DoRegister(string, string) error {
+       return nil
+}
+
+func (f *baseRegistryTestFacade) DoUnregister(string, string) error {
+       return nil
+}
+
+func (f *baseRegistryTestFacade) CloseAndNilClient() {}
+
+func (f *baseRegistryTestFacade) CloseListener() {
+       if f.listener != nil {
+               f.listener.Close()
+       }
+}
+
+func (f *baseRegistryTestFacade) InitListeners() {}
+
+type baseRegistryTestListener struct {
+       closed    chan struct{}
+       closeOnce sync.Once
+}
+
+func (l *baseRegistryTestListener) Next() (*ServiceEvent, error) {
+       <-l.closed
+       return nil, errors.New("listener closed")
+}
+
+func (l *baseRegistryTestListener) Close() {
+       l.closeOnce.Do(func() { close(l.closed) })
+}
+
+type gatedRegistryTestListener struct {
+       release   chan struct{}
+       nextReady chan struct{}
+       closeOnce sync.Once
+}
+
+func (l *gatedRegistryTestListener) Next() (*ServiceEvent, error) {
+       close(l.nextReady)
+       <-l.release
+       return &ServiceEvent{}, nil
+}
+
+func (l *gatedRegistryTestListener) Close() {
+       l.closeOnce.Do(func() {})
+}
+
+type baseRegistryTestNotify struct {
+       notified atomic.Int32
+}
+
+func (n *baseRegistryTestNotify) Notify(*ServiceEvent) {
+       n.notified.Add(1)
+}
+
+func (*baseRegistryTestNotify) NotifyAll([]*ServiceEvent, func()) {}
+
+func TestBaseRegistrySubscribeDestroyInterruptsRetryDelay(t *testing.T) {
+       listener := &baseRegistryTestListener{closed: make(chan struct{})}
+       facade := newBaseRegistryTestFacade(listener, errors.New("subscribe 
failed"))
+
+       subscribeDone := make(chan error, 1)
+       go func() {
+               subscribeDone <- facade.Subscribe(common.NewURLWithOptions(), 
&baseRegistryTestNotify{})
+       }()
+
+       select {
+       case <-facade.firstSubscribe:
+       case <-time.After(time.Second):
+               t.Fatal("Subscribe did not attempt the initial subscription")
+       }
+
+       start := time.Now()
+       facade.Destroy()
+
+       select {
+       case err := <-subscribeDone:
+               require.Error(t, err)
+       case <-time.After(time.Second):
+               t.Fatal("Subscribe remained blocked after Destroy")
+       }
+       require.Less(t, time.Since(start), time.Second)
+}
+
+func TestBaseRegistrySubscribeReturnsUnavailableAfterDestroyDuringSubscribe(t 
*testing.T) {
+       listener := &baseRegistryTestListener{closed: make(chan struct{})}
+       facade := newBaseRegistryTestFacade(listener, errors.New("subscribe 
failed"))
+       facade.subscribeRelease = make(chan struct{})
+
+       subscribeDone := make(chan error, 1)
+       go func() {
+               subscribeDone <- facade.Subscribe(common.NewURLWithOptions(), 
&baseRegistryTestNotify{})
+       }()
+
+       select {
+       case <-facade.firstSubscribe:
+       case <-time.After(time.Second):
+               t.Fatal("Subscribe did not attempt the initial subscription")
+       }
+       facade.Destroy()
+       close(facade.subscribeRelease)
+
+       require.ErrorIs(t, <-subscribeDone, errBaseRegistryUnavailable)
+}
+
+func TestBaseRegistrySubscribeDoesNotNotifyAfterDestroy(t *testing.T) {
+       listener := &gatedRegistryTestListener{
+               release:   make(chan struct{}),
+               nextReady: make(chan struct{}),
+       }
+       facade := newBaseRegistryTestFacade(listener, nil)
+       notify := &baseRegistryTestNotify{}
+       subscribeDone := make(chan error, 1)
+       go func() {
+               subscribeDone <- facade.Subscribe(common.NewURLWithOptions(), 
notify)
+       }()
+
+       select {
+       case <-listener.nextReady:
+       case <-time.After(time.Second):
+               t.Fatal("Subscribe did not enter listener.Next")
+       }
+
+       facade.Destroy()
+       close(listener.release)
+
+       select {
+       case err := <-subscribeDone:
+               require.Error(t, err)
+       case <-time.After(time.Second):
+               t.Fatal("Subscribe did not exit after Destroy")
+       }
+       require.Zero(t, notify.notified.Load())
+}
diff --git a/remoting/getty/getty_client.go b/remoting/getty/getty_client.go
index 77d8821a7..8fd452101 100644
--- a/remoting/getty/getty_client.go
+++ b/remoting/getty/getty_client.go
@@ -145,16 +145,17 @@ type Options struct {
 
 // Client : some configuration for network communication.
 type Client struct {
-       addr               string
-       opts               Options
-       conf               ClientConfig
-       mux                sync.RWMutex
-       sslEnabled         bool
-       clientClosed       bool
-       gettyClient        *gettyRPCClient
-       gettyClientMux     sync.RWMutex
-       gettyClientCreated atomic.Bool
-       codec              remoting.Codec
+       addr           string
+       opts           Options
+       conf           ClientConfig
+       connectMu      sync.Mutex
+       closeOnce      sync.Once
+       done           chan struct{}
+       sslEnabled     bool
+       closed         atomic.Bool
+       gettyClient    *gettyRPCClient
+       gettyClientMux sync.RWMutex
+       codec          remoting.Codec
 }
 
 // NewClient create client
@@ -168,10 +169,9 @@ func NewClient(opt Options) *Client {
        }
 
        c := &Client{
-               opts:         opt,
-               clientClosed: false,
+               opts: opt,
+               done: make(chan struct{}),
        }
-       c.gettyClientCreated.Store(false)
        return c
 }
 
@@ -180,6 +180,9 @@ func (c *Client) SetExchangeClient(client 
*remoting.ExchangeClient) {
 
 // Connect init client and try to connection.
 func (c *Client) Connect(url *common.URL) error {
+       if c.closed.Load() {
+               return errClientClosed
+       }
        initClient(url)
        c.conf = *clientConf
        c.sslEnabled = c.conf.SSLEnabled
@@ -195,14 +198,19 @@ func (c *Client) Connect(url *common.URL) error {
 
 // Close close network connection
 func (c *Client) Close() {
-       c.mux.Lock()
-       client := c.gettyClient
-       c.gettyClient = nil
-       c.clientClosed = true
-       c.mux.Unlock()
-       if client != nil {
-               client.close()
-       }
+       c.closeOnce.Do(func() {
+               c.closed.Store(true)
+               if c.done != nil {
+                       close(c.done)
+               }
+               c.gettyClientMux.Lock()
+               client := c.gettyClient
+               c.gettyClient = nil
+               c.gettyClientMux.Unlock()
+               if client != nil {
+                       client.close()
+               }
+       })
 }
 
 // Request send request
@@ -252,34 +260,55 @@ func (c *Client) IsAvailable() bool {
 }
 
 func (c *Client) selectSession(addr string) (*gettyRPCClient, getty.Session, 
error) {
-       c.mux.RLock()
-       defer c.mux.RUnlock()
-       if c.clientClosed {
-               return nil, nil, perrors.New("client have been closed")
+       if c.closed.Load() {
+               return nil, nil, errClientClosed
        }
 
-       if !c.gettyClientCreated.Load() {
-               c.gettyClientMux.Lock()
-               if c.gettyClient == nil {
-                       rpcClientConn, rpcErr := newGettyRPCClientConn(c, addr)
-                       if rpcErr != nil {
-                               c.gettyClientMux.Unlock()
-                               return nil, nil, perrors.WithStack(rpcErr)
-                       }
-                       c.gettyClientCreated.Store(true)
-                       c.gettyClient = rpcClientConn
+       c.gettyClientMux.RLock()
+       client := c.gettyClient
+       c.gettyClientMux.RUnlock()
+       if client == nil {
+               var err error
+               client, err = c.getOrCreateGettyClient(addr, 
newGettyRPCClientConn)
+               if err != nil {
+                       return nil, nil, perrors.WithStack(err)
                }
-               client := c.gettyClient
-               session := c.gettyClient.selectSession()
-               c.gettyClientMux.Unlock()
-               return client, session, nil
+       }
+
+       if c.closed.Load() {
+               return nil, nil, errClientClosed
+       }
+       return client, client.selectSession(), nil
+}
+
+func (c *Client) getOrCreateGettyClient(addr string, newClientConn 
func(*Client, string) (*gettyRPCClient, error)) (*gettyRPCClient, error) {
+       c.connectMu.Lock()
+       defer c.connectMu.Unlock()
+
+       if c.closed.Load() {
+               return nil, errClientClosed
        }
        c.gettyClientMux.RLock()
        client := c.gettyClient
-       session := c.gettyClient.selectSession()
        c.gettyClientMux.RUnlock()
-       return client, session, nil
+       if client != nil {
+               return client, nil
+       }
+
+       client, err := newClientConn(c, addr)
+       if err != nil {
+               return nil, err
+       }
 
+       c.gettyClientMux.Lock()
+       if c.closed.Load() {
+               c.gettyClientMux.Unlock()
+               _ = client.close()
+               return nil, errClientClosed
+       }
+       c.gettyClient = client
+       c.gettyClientMux.Unlock()
+       return client, nil
 }
 
 func (c *Client) transfer(session getty.Session, request *remoting.Request, 
timeout time.Duration) (int, int, error) {
@@ -287,10 +316,14 @@ func (c *Client) transfer(session getty.Session, request 
*remoting.Request, time
        return totalLen, sendLen, perrors.WithStack(err)
 }
 
-func (c *Client) resetRpcConn() {
+func (c *Client) resetRpcConn(expected *gettyRPCClient) {
+       c.connectMu.Lock()
+       defer c.connectMu.Unlock()
+
        c.gettyClientMux.Lock()
+       defer c.gettyClientMux.Unlock()
+       if c.gettyClient != expected {
+               return
+       }
        c.gettyClient = nil
-       c.gettyClientCreated.Store(false)
-       c.gettyClientMux.Unlock()
-
 }
diff --git a/remoting/getty/getty_client_test.go 
b/remoting/getty/getty_client_test.go
index 0df6b65dc..6e15c396e 100644
--- a/remoting/getty/getty_client_test.go
+++ b/remoting/getty/getty_client_test.go
@@ -20,6 +20,7 @@ package getty
 import (
        "bytes"
        "context"
+       "net"
        "reflect"
        "sync"
        "testing"
@@ -27,6 +28,8 @@ import (
 )
 
 import (
+       dubboGetty "github.com/apache/dubbo-getty"
+
        hessian "github.com/apache/dubbo-go-hessian2"
 
        perrors "github.com/pkg/errors"
@@ -46,6 +49,16 @@ import (
        "dubbo.apache.org/dubbo-go/v3/remoting"
 )
 
+type closeTrackingGettyClient struct {
+       dubboGetty.Client
+       closeOnce sync.Once
+       closed    chan struct{}
+}
+
+func (c *closeTrackingGettyClient) Close() {
+       c.closeOnce.Do(func() { close(c.closed) })
+}
+
 func TestRunSuite(t *testing.T) {
        svr, url := InitTest(t)
        client := getClient(url)
@@ -331,3 +344,134 @@ func TestInitClientTLS(t *testing.T) {
                assert.False(t, clientConf.SSLEnabled)
        })
 }
+
+func TestGettyConnectWaitStopsWhenClosed(t *testing.T) {
+       client := NewClient(Options{ConnectTimeout: 5 * time.Second})
+       started := make(chan struct{})
+       var startOnce sync.Once
+       available := func() bool {
+               startOnce.Do(func() { close(started) })
+               return false
+       }
+       waitDone := make(chan error, 1)
+       go func() {
+               waitDone <- waitForGettyClient("127.0.0.1:1", 
client.opts.ConnectTimeout, available, client.done)
+       }()
+       select {
+       case <-started:
+       case <-time.After(time.Second):
+               t.Fatal("connection wait did not start")
+       }
+
+       start := time.Now()
+       client.Close()
+       err := <-waitDone
+
+       require.Error(t, err)
+       require.ErrorIs(t, err, errClientClosed)
+       require.Less(t, time.Since(start), time.Second)
+}
+
+func TestGettyCloseAfterConnectionReadyBeforePublish(t *testing.T) {
+       client := NewClient(Options{ConnectTimeout: time.Second})
+       fakeGettyClient := &closeTrackingGettyClient{closed: make(chan 
struct{})}
+       fakeRPCClient := &gettyRPCClient{gettyClient: fakeGettyClient}
+       factoryReady := make(chan struct{})
+       releaseFactory := make(chan struct{})
+       connectDone := make(chan error, 1)
+
+       go func() {
+               _, err := client.getOrCreateGettyClient("", func(_ *Client, _ 
string) (*gettyRPCClient, error) {
+                       close(factoryReady)
+                       <-releaseFactory
+                       return fakeRPCClient, nil
+               })
+               connectDone <- err
+       }()
+
+       select {
+       case <-factoryReady:
+       case <-time.After(time.Second):
+               t.Fatal("connection factory did not become ready")
+       }
+
+       client.Close()
+       close(releaseFactory)
+
+       select {
+       case err := <-connectDone:
+               require.ErrorIs(t, err, errClientClosed)
+       case <-time.After(time.Second):
+               t.Fatal("connection creation did not finish after release")
+       }
+
+       select {
+       case <-fakeGettyClient.closed:
+       case <-time.After(time.Second):
+               t.Fatal("unpublished connection was not closed")
+       }
+
+       client.gettyClientMux.RLock()
+       require.Nil(t, client.gettyClient)
+       client.gettyClientMux.RUnlock()
+}
+
+func TestGettyConnectWaitHonorsTimeout(t *testing.T) {
+       start := time.Now()
+       err := waitForGettyClient("127.0.0.1:1", 30*time.Millisecond,
+               func() bool { return false },
+               nil,
+       )
+
+       require.Error(t, err)
+       require.NotErrorIs(t, err, errClientClosed)
+       require.Less(t, time.Since(start), time.Second)
+}
+
+func TestGettyNewConnectionStopsWhenClientCloses(t *testing.T) {
+       listener, err := net.Listen("tcp", "127.0.0.1:0")
+       require.NoError(t, err)
+       addr := listener.Addr().String()
+       require.NoError(t, listener.Close())
+
+       client := NewClient(Options{ConnectTimeout: 5 * time.Second})
+       client.conf = *GetDefaultClientConfig()
+       connectDone := make(chan error, 1)
+       go func() {
+               _, connectErr := newGettyRPCClientConn(client, addr)
+               connectDone <- connectErr
+       }()
+       time.AfterFunc(20*time.Millisecond, client.Close)
+
+       start := time.Now()
+       select {
+       case err := <-connectDone:
+               require.Error(t, err)
+               require.ErrorIs(t, err, errClientClosed)
+       case <-time.After(time.Second):
+               t.Fatal("newGettyRPCClientConn remained blocked after Close")
+       }
+       require.Less(t, time.Since(start), time.Second)
+}
+
+func TestClientCloseDoesNotWaitForConnectLock(t *testing.T) {
+       client := NewClient(Options{ConnectTimeout: time.Second, 
RequestTimeout: time.Second})
+       client.connectMu.Lock()
+       closeDone := make(chan struct{})
+       go func() {
+               client.Close()
+               close(closeDone)
+       }()
+
+       select {
+       case <-closeDone:
+       case <-time.After(time.Second):
+               t.Fatal("Client.Close waited for the connection lock")
+       }
+       client.connectMu.Unlock()
+
+       require.True(t, client.closed.Load())
+       _, _, err := client.selectSession("")
+       require.Error(t, err)
+       require.ErrorIs(t, err, errClientClosed)
+}
diff --git a/remoting/getty/pool.go b/remoting/getty/pool.go
index 2c70e415a..7434e11df 100644
--- a/remoting/getty/pool.go
+++ b/remoting/getty/pool.go
@@ -75,25 +75,10 @@ func newGettyRPCClientConn(rpcClient *Client, addr string) 
(*gettyRPCClient, err
        }
        go c.gettyClient.RunEventLoop(c.newSession)
 
-       idx := 1
-       start := time.Now()
        connectTimeout := rpcClient.opts.ConnectTimeout
-       for {
-               idx++
-               if c.isAvailable() {
-                       break
-               }
-
-               if time.Since(start) > connectTimeout {
-                       c.gettyClient.Close()
-                       return nil, perrors.New(fmt.Sprintf("failed to create 
client connection to %s in %s", addr, connectTimeout))
-               }
-
-               interval := time.Millisecond * time.Duration(idx)
-               if interval > time.Duration(100e6) {
-                       interval = 100e6 // 100 ms
-               }
-               time.Sleep(interval)
+       if err := waitForGettyClient(addr, connectTimeout, c.isAvailable, 
rpcClient.done); err != nil {
+               c.gettyClient.Close()
+               return nil, err
        }
        logger.Debug("[Remoting][Getty] client init ok")
        c.updateActive(time.Now().Unix())
@@ -101,6 +86,32 @@ func newGettyRPCClientConn(rpcClient *Client, addr string) 
(*gettyRPCClient, err
        return c, nil
 }
 
+func waitForGettyClient(addr string, timeout time.Duration, available func() 
bool, done <-chan struct{}) error {
+       start := time.Now()
+       for idx := 2; ; idx++ {
+               if available() {
+                       return nil
+               }
+               if time.Since(start) > timeout {
+                       return perrors.New(fmt.Sprintf("failed to create client 
connection to %s in %s", addr, timeout))
+               }
+
+               interval := min(time.Millisecond*time.Duration(idx), 
100*time.Millisecond)
+               timer := time.NewTimer(interval)
+               select {
+               case <-timer.C:
+               case <-done:
+                       if !timer.Stop() {
+                               select {
+                               case <-timer.C:
+                               default:
+                               }
+                       }
+                       return errClientClosed
+               }
+       }
+}
+
 func (c *gettyRPCClient) updateActive(active int64) {
        c.active.Store(active)
 }
@@ -220,7 +231,7 @@ func (c *gettyRPCClient) removeSession(session 
getty.Session) {
                }
        }()
        if removeFlag {
-               c.rpcClient.resetRpcConn()
+               c.rpcClient.resetRpcConn(c)
                c.close()
        }
 }

Reply via email to