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

Alanxtl 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 ca39fe208 fix(graceful_shutdown, loadbalance): prevent hard exits and 
P2C panics (#3591)
ca39fe208 is described below

commit ca39fe208417ffb165e024d8a79948947411007d
Author: Nene7ko_ <[email protected]>
AuthorDate: Thu Aug 20 13:35:36 2026 +0800

    fix(graceful_shutdown, loadbalance): prevent hard exits and P2C panics 
(#3591)
    
    * fix crash risks in shutdown and p2c load balancer
    
    * fix graceful shutdown timeout propagation
    
    * test: fix graceful shutdown helper arguments
    
    * test: avoid leaking graceful shutdown signal goroutine
    
    * test: isolate P2C metrics state
---
 cluster/loadbalance/p2c/loadbalance.go      |  22 +++--
 cluster/loadbalance/p2c/loadbalance_test.go | 136 +++++++++++++++++++++++++++-
 graceful_shutdown/shutdown.go               |  44 ++++++---
 graceful_shutdown/shutdown_test.go          |  92 ++++++++++++++++++-
 server/server.go                            |  10 +-
 server/server_test.go                       |  47 ++++++++++
 6 files changed, 325 insertions(+), 26 deletions(-)

diff --git a/cluster/loadbalance/p2c/loadbalance.go 
b/cluster/loadbalance/p2c/loadbalance.go
index d5d82f000..cca68494e 100644
--- a/cluster/loadbalance/p2c/loadbalance.go
+++ b/cluster/loadbalance/p2c/loadbalance.go
@@ -19,7 +19,6 @@ package p2c
 
 import (
        "errors"
-       "fmt"
        "math/rand"
        "sync"
        "time"
@@ -127,6 +126,10 @@ func (l *p2cLoadBalance) Select(invokers []base.Invoker, 
invocation base.Invocat
                logger.Warnf("[Loadbalance][P2C] get method metrics err=%v", 
err)
                return nil
        }
+       if remainingIIface == nil {
+               logger.Debugf("[Loadbalance][P2C] select invoker i=%d (nil 
metrics)", i)
+               return invokers[i]
+       }
 
        // TODO(justxuewei): It should have a strategy to drop some metrics 
after a period of time.
        remainingJIface, err := m.GetMethodMetrics(invokers[j].GetURL(), 
methodName, metrics.HillClimbing)
@@ -138,18 +141,25 @@ func (l *p2cLoadBalance) Select(invokers []base.Invoker, 
invocation base.Invocat
                logger.Warnf("[Loadbalance][P2C] get method metrics err=%v", 
err)
                return nil
        }
+       if remainingJIface == nil {
+               logger.Debugf("[Loadbalance][P2C] select invoker j=%d (nil 
metrics)", j)
+               return invokers[j]
+       }
 
-       // Convert interface to int, if the type is unexpected, panic 
immediately
+       // Convert interface to uint64. If one value has an unexpected type, 
prefer that
+       // invoker so the caller can collect fresh metrics instead of crashing.
        remainingI, ok := remainingIIface.(uint64)
        if !ok {
-               panic(fmt.Sprintf("[Loadbalance][P2C] type check failed: key=%s 
expected=uint64 actual=%T",
-                       metrics.HillClimbing, remainingIIface))
+               logger.Warnf("[Loadbalance][P2C] type check failed: key=%s 
expected=uint64 actual=%T",
+                       metrics.HillClimbing, remainingIIface)
+               return invokers[i]
        }
 
        remainingJ, ok := remainingJIface.(uint64)
        if !ok {
-               panic(fmt.Sprintf("[Loadbalance][P2C] type check failed: key=%s 
expected=uint64 actual=%T",
-                       metrics.HillClimbing, remainingJIface))
+               logger.Warnf("[Loadbalance][P2C] type check failed: key=%s 
expected=uint64 actual=%T",
+                       metrics.HillClimbing, remainingJIface)
+               return invokers[j]
        }
 
        logger.Debugf("[Loadbalance][P2C] compare remaining capacity i=%d 
val=%d j=%d val=%d", i, remainingI, j, remainingJ)
diff --git a/cluster/loadbalance/p2c/loadbalance_test.go 
b/cluster/loadbalance/p2c/loadbalance_test.go
index e8ac720b1..5d39b17dd 100644
--- a/cluster/loadbalance/p2c/loadbalance_test.go
+++ b/cluster/loadbalance/p2c/loadbalance_test.go
@@ -56,6 +56,16 @@ func TestDefaultRnd(t *testing.T) {
        })
 }
 
+func setLocalMetrics(t *testing.T, localMetrics metrics.Metrics) {
+       t.Helper()
+
+       originalLocalMetrics := metrics.LocalMetrics
+       metrics.LocalMetrics = localMetrics
+       t.Cleanup(func() {
+               metrics.LocalMetrics = originalLocalMetrics
+       })
+}
+
 func TestLoadBalance(t *testing.T) {
        // Create P2C load balancer with deterministic randomPicker for 
repeatable tests.
        // Always returns fixed indices (0,1) except when n <= 1.
@@ -90,7 +100,7 @@ func TestLoadBalance(t *testing.T) {
                defer ctrl.Finish()
 
                m := metrics.NewMockMetrics(ctrl)
-               metrics.LocalMetrics = m
+               setLocalMetrics(t, m)
 
                url0, _ := 
common.NewURL("dubbo://192.168.1.0:20000/com.ikurento.user.UserProvider")
                url1, _ := 
common.NewURL("dubbo://192.168.1.1:20000/com.ikurento.user.UserProvider")
@@ -119,7 +129,7 @@ func TestLoadBalance(t *testing.T) {
                defer ctrl.Finish()
 
                m := metrics.NewMockMetrics(ctrl)
-               metrics.LocalMetrics = m
+               setLocalMetrics(t, m)
 
                url0, _ := 
common.NewURL("dubbo://192.168.1.0:20000/com.ikurento.user.UserProvider")
                url1, _ := 
common.NewURL("dubbo://192.168.1.1:20000/com.ikurento.user.UserProvider")
@@ -150,7 +160,7 @@ func TestLoadBalance(t *testing.T) {
                defer ctrl.Finish()
 
                m := metrics.NewMockMetrics(ctrl)
-               metrics.LocalMetrics = m
+               setLocalMetrics(t, m)
 
                url0, _ := 
common.NewURL("dubbo://192.168.1.0:20000/com.ikurento.user.UserProvider")
                url1, _ := 
common.NewURL("dubbo://192.168.1.1:20000/com.ikurento.user.UserProvider")
@@ -177,7 +187,7 @@ func TestLoadBalance(t *testing.T) {
                defer ctrl.Finish()
 
                m := metrics.NewMockMetrics(ctrl)
-               metrics.LocalMetrics = m
+               setLocalMetrics(t, m)
 
                url0, _ := 
common.NewURL("dubbo://192.168.1.0:20000/com.ikurento.user.UserProvider")
                url1, _ := 
common.NewURL("dubbo://192.168.1.1:20000/com.ikurento.user.UserProvider")
@@ -204,4 +214,122 @@ func TestLoadBalance(t *testing.T) {
                assert.Equal(t, ivkArr[1].GetURL().String(), 
ivk.GetURL().String())
        })
 
+       t.Run("metrics i nil", func(t *testing.T) {
+               ctrl := gomock.NewController(t)
+               defer ctrl.Finish()
+
+               m := metrics.NewMockMetrics(ctrl)
+               setLocalMetrics(t, m)
+
+               url0, _ := 
common.NewURL("dubbo://192.168.1.0:20000/com.ikurento.user.UserProvider")
+               url1, _ := 
common.NewURL("dubbo://192.168.1.1:20000/com.ikurento.user.UserProvider")
+
+               m.EXPECT().
+                       GetMethodMetrics(gomock.Eq(url0), 
gomock.Eq(invocation.MethodName()), gomock.Eq(metrics.HillClimbing)).
+                       Times(1).
+                       Return(nil, nil)
+
+               ivkArr := []base.Invoker{
+                       base.NewBaseInvoker(url0),
+                       base.NewBaseInvoker(url1),
+               }
+
+               assert.NotPanics(t, func() {
+                       ivk := lb.Select(ivkArr, invocation)
+                       assert.Equal(t, ivkArr[0].GetURL().String(), 
ivk.GetURL().String())
+               })
+       })
+
+       t.Run("metrics j nil", func(t *testing.T) {
+               ctrl := gomock.NewController(t)
+               defer ctrl.Finish()
+
+               m := metrics.NewMockMetrics(ctrl)
+               setLocalMetrics(t, m)
+
+               url0, _ := 
common.NewURL("dubbo://192.168.1.0:20000/com.ikurento.user.UserProvider")
+               url1, _ := 
common.NewURL("dubbo://192.168.1.1:20000/com.ikurento.user.UserProvider")
+
+               m.EXPECT().
+                       GetMethodMetrics(gomock.Eq(url0), 
gomock.Eq(invocation.MethodName()), gomock.Eq(metrics.HillClimbing)).
+                       Times(1).
+                       Return(uint64(0), nil)
+
+               m.EXPECT().
+                       GetMethodMetrics(gomock.Eq(url1), 
gomock.Eq(invocation.MethodName()), gomock.Eq(metrics.HillClimbing)).
+                       Times(1).
+                       Return(nil, nil)
+
+               ivkArr := []base.Invoker{
+                       base.NewBaseInvoker(url0),
+                       base.NewBaseInvoker(url1),
+               }
+
+               assert.NotPanics(t, func() {
+                       ivk := lb.Select(ivkArr, invocation)
+                       assert.Equal(t, ivkArr[1].GetURL().String(), 
ivk.GetURL().String())
+               })
+       })
+
+       t.Run("metrics i wrong type", func(t *testing.T) {
+               ctrl := gomock.NewController(t)
+               defer ctrl.Finish()
+
+               m := metrics.NewMockMetrics(ctrl)
+               setLocalMetrics(t, m)
+
+               url0, _ := 
common.NewURL("dubbo://192.168.1.0:20000/com.ikurento.user.UserProvider")
+               url1, _ := 
common.NewURL("dubbo://192.168.1.1:20000/com.ikurento.user.UserProvider")
+
+               m.EXPECT().
+                       GetMethodMetrics(gomock.Eq(url0), 
gomock.Eq(invocation.MethodName()), gomock.Eq(metrics.HillClimbing)).
+                       Times(1).
+                       Return("bad-metrics", nil)
+
+               m.EXPECT().
+                       GetMethodMetrics(gomock.Eq(url1), 
gomock.Eq(invocation.MethodName()), gomock.Eq(metrics.HillClimbing)).
+                       Times(1).
+                       Return(uint64(10), nil)
+
+               ivkArr := []base.Invoker{
+                       base.NewBaseInvoker(url0),
+                       base.NewBaseInvoker(url1),
+               }
+
+               assert.NotPanics(t, func() {
+                       ivk := lb.Select(ivkArr, invocation)
+                       assert.Equal(t, ivkArr[0].GetURL().String(), 
ivk.GetURL().String())
+               })
+       })
+
+       t.Run("metrics j wrong type", func(t *testing.T) {
+               ctrl := gomock.NewController(t)
+               defer ctrl.Finish()
+
+               m := metrics.NewMockMetrics(ctrl)
+               setLocalMetrics(t, m)
+
+               url0, _ := 
common.NewURL("dubbo://192.168.1.0:20000/com.ikurento.user.UserProvider")
+               url1, _ := 
common.NewURL("dubbo://192.168.1.1:20000/com.ikurento.user.UserProvider")
+
+               m.EXPECT().
+                       GetMethodMetrics(gomock.Eq(url0), 
gomock.Eq(invocation.MethodName()), gomock.Eq(metrics.HillClimbing)).
+                       Times(1).
+                       Return(uint64(10), nil)
+
+               m.EXPECT().
+                       GetMethodMetrics(gomock.Eq(url1), 
gomock.Eq(invocation.MethodName()), gomock.Eq(metrics.HillClimbing)).
+                       Times(1).
+                       Return("bad-metrics", nil)
+
+               ivkArr := []base.Invoker{
+                       base.NewBaseInvoker(url0),
+                       base.NewBaseInvoker(url1),
+               }
+
+               assert.NotPanics(t, func() {
+                       ivk := lb.Select(ivkArr, invocation)
+                       assert.Equal(t, ivkArr[1].GetURL().String(), 
ivk.GetURL().String())
+               })
+       })
 }
diff --git a/graceful_shutdown/shutdown.go b/graceful_shutdown/shutdown.go
index b024a1f29..bd5eb14b1 100644
--- a/graceful_shutdown/shutdown.go
+++ b/graceful_shutdown/shutdown.go
@@ -64,12 +64,14 @@ var (
        shutdownConfigMu sync.RWMutex
        shutdownConfig   *global.ShutdownConfig
 
-       shutdownOnce    sync.Once
-       shutdownStarted atomic.Bool
-       shutdownDone    = make(chan struct{})
-       shutdownResult  error
+       shutdownOnce        sync.Once
+       shutdownStarted     atomic.Bool
+       shutdownDone        = make(chan struct{})
+       shutdownResult      error
+       shutdownSignalError = make(chan error, 1)
 
        signalNotify = signal.Notify
+       signalStop   = signal.Stop
 )
 
 type shutdownConfigSetter interface {
@@ -108,23 +110,23 @@ func Init(opts ...Option) {
                        signalNotify(signals, ShutdownSignals...)
 
                        go func() {
+                               defer signalStop(signals)
+
                                sig := <-signals
                                logger.Infof("[GracefulShutdown] get signal %s, 
applicationConfig will shutdown.", sig)
-                               // fallback timeout
-                               time.AfterFunc(totalTimeout(newOpts.Shutdown), 
func() {
-                                       logger.Warn("[GracefulShutdown] 
shutdown gracefully timeout, applicationConfig will shutdown immediately. ")
-                                       os.Exit(0)
-                               })
-                               if err := Shutdown(context.Background()); err 
!= nil {
+                               ctx, cancel := 
context.WithTimeout(context.Background(), totalTimeout(newOpts.Shutdown))
+                               defer cancel()
+
+                               if err := Shutdown(ctx); err != nil {
                                        logger.Warnf("[GracefulShutdown] 
shutdown completed, err=%v", err)
+                                       reportShutdownError(err)
                                }
-                               // those signals' original behavior is exit 
with dump ths stack, so we try to keep the behavior
+                               // Preserve heap-dump behavior for dump signals 
without terminating the host process.
                                for _, dumpSignal := range 
DumpHeapShutdownSignals {
                                        if sig == dumpSignal {
                                                
debug.WriteHeapDump(os.Stdout.Fd())
                                        }
                                }
-                               os.Exit(0)
                        }()
                }
        })
@@ -134,6 +136,24 @@ func Done() <-chan struct{} {
        return shutdownDone
 }
 
+// ShutdownError returns the error channel for failures encountered during an
+// internal signal-triggered shutdown.
+func ShutdownError() <-chan error {
+       return shutdownSignalError
+}
+
+func reportShutdownError(err error) {
+       if err == nil {
+               return
+       }
+
+       select {
+       case shutdownSignalError <- err:
+       default:
+               logger.Warnf("[GracefulShutdown] shutdown error channel is 
full, err=%v", err)
+       }
+}
+
 func IsDone() bool {
        select {
        case <-shutdownDone:
diff --git a/graceful_shutdown/shutdown_test.go 
b/graceful_shutdown/shutdown_test.go
index 63fe68c2a..fb9aa1d71 100644
--- a/graceful_shutdown/shutdown_test.go
+++ b/graceful_shutdown/shutdown_test.go
@@ -20,7 +20,9 @@ package graceful_shutdown
 import (
        "context"
        "errors"
+       "fmt"
        "os"
+       "os/exec"
        "os/signal"
        "sync"
        "sync/atomic"
@@ -114,7 +116,9 @@ func resetShutdownTestState() {
        shutdownStarted = atomic.Bool{}
        shutdownDone = make(chan struct{})
        shutdownResult = nil
+       shutdownSignalError = make(chan error, 1)
        signalNotify = signal.Notify
+       signalStop = signal.Stop
 }
 
 func TestInit(t *testing.T) {
@@ -138,11 +142,11 @@ func TestInit(t *testing.T) {
        })
 
        // Test with default options
-       Init()
+       Init(WithoutInternalSignal())
 
        // Test with custom options
        customTimeout := 120 * time.Second
-       Init(WithTimeout(customTimeout))
+       Init(WithTimeout(customTimeout), WithoutInternalSignal())
 
        // Remove mock filters
        extension.UnregisterFilter(constant.GracefulShutdownConsumerFilterKey)
@@ -175,6 +179,90 @@ func TestInitReturnsWhenGracefulShutdownFilterMissing(t 
*testing.T) {
        assert.Nil(t, shutdownConfig)
 }
 
+func TestInitInternalSignalTriggersShutdownWithoutProcessExit(t *testing.T) {
+       if os.Getenv("DUBBO_GO_GRACEFUL_SHUTDOWN_HELPER") == "1" {
+               runInternalSignalShutdownTest(t)
+               fmt.Fprintln(os.Stdout, "graceful shutdown helper completed")
+               return
+       }
+
+       cmd := exec.Command(os.Args[0], 
"-test.run=^TestInitInternalSignalTriggersShutdownWithoutProcessExit$")
+       cmd.Env = append(os.Environ(), "DUBBO_GO_GRACEFUL_SHUTDOWN_HELPER=1")
+       output, err := cmd.CombinedOutput()
+       require.NoError(t, err, string(output))
+       assert.Contains(t, string(output), "graceful shutdown helper completed")
+}
+
+func runInternalSignalShutdownTest(t *testing.T) {
+       resetShutdownTestState()
+
+       mockConsumerFilter := &MockFilter{}
+       mockProviderFilter := &MockFilter{}
+       mockConsumerFilter.On("Set", mock.Anything, mock.Anything).Return()
+       mockProviderFilter.On("Set", mock.Anything, mock.Anything).Return()
+
+       extension.SetFilter(constant.GracefulShutdownConsumerFilterKey, func() 
filter.Filter {
+               return mockConsumerFilter
+       })
+       extension.SetFilter(constant.GracefulShutdownProviderFilterKey, func() 
filter.Filter {
+               return mockProviderFilter
+       })
+       t.Cleanup(func() {
+               
extension.UnregisterFilter(constant.GracefulShutdownConsumerFilterKey)
+               
extension.UnregisterFilter(constant.GracefulShutdownProviderFilterKey)
+               resetShutdownTestState()
+       })
+
+       signalNotify = func(signals chan<- os.Signal, _ ...os.Signal) {
+               signals <- os.Interrupt
+       }
+       stopCalled := atomic.Bool{}
+       signalStop = func(chan<- os.Signal) {
+               stopCalled.Store(true)
+       }
+
+       cfg := global.DefaultShutdownConfig()
+       internalSignal := true
+       cfg.InternalSignal = &internalSignal
+       cfg.ConsumerUpdateWaitTime = "0s"
+       cfg.StepTimeout = "0s"
+       cfg.NotifyTimeout = "10ms"
+       cfg.OfflineRequestWindowTimeout = "0s"
+
+       Init(SetShutdownConfig(cfg))
+
+       select {
+       case <-Done():
+       case <-time.After(time.Second):
+               t.Fatal("shutdown did not complete after internal signal")
+       }
+
+       require.NoError(t, Shutdown(context.Background()))
+       require.Eventually(t, func() bool {
+               return stopCalled.Load()
+       }, time.Second, time.Millisecond)
+
+       select {
+       case err := <-ShutdownError():
+               t.Fatalf("unexpected shutdown error: %v", err)
+       default:
+       }
+}
+
+func TestReportShutdownError(t *testing.T) {
+       resetShutdownTestState()
+
+       expectedErr := context.DeadlineExceeded
+       reportShutdownError(expectedErr)
+
+       select {
+       case err := <-ShutdownError():
+               require.ErrorIs(t, err, expectedErr)
+       case <-time.After(time.Second):
+               t.Fatal("shutdown error was not reported")
+       }
+}
+
 func TestShutdownClosesDoneAndRunsOnce(t *testing.T) {
        resetShutdownTestState()
 
diff --git a/server/server.go b/server/server.go
index 9969605f4..d1da95b69 100644
--- a/server/server.go
+++ b/server/server.go
@@ -412,13 +412,19 @@ func (s *Server) ServeContext(ctx context.Context) error {
                select {
                case <-graceful_shutdown.Done():
                        return graceful_shutdown.Shutdown(context.Background())
+               case err := <-graceful_shutdown.ShutdownError():
+                       return err
                case <-done:
                        return graceful_shutdown.Shutdown(ctx)
                }
        }
 
-       <-graceful_shutdown.Done()
-       return graceful_shutdown.Shutdown(context.Background())
+       select {
+       case <-graceful_shutdown.Done():
+               return graceful_shutdown.Shutdown(context.Background())
+       case err := <-graceful_shutdown.ShutdownError():
+               return err
+       }
 }
 
 func (s *Server) rollbackServeStartWithCause(cause error, 
serviceInstanceRegistered bool) error {
diff --git a/server/server_test.go b/server/server_test.go
index 6217d4c5f..1b6222997 100644
--- a/server/server_test.go
+++ b/server/server_test.go
@@ -83,6 +83,9 @@ var gracefulShutdownDone chan struct{}
 //go:linkname gracefulShutdownResult 
dubbo.apache.org/dubbo-go/v3/graceful_shutdown.shutdownResult
 var gracefulShutdownResult error
 
+//go:linkname gracefulShutdownSignalError 
dubbo.apache.org/dubbo-go/v3/graceful_shutdown.shutdownSignalError
+var gracefulShutdownSignalError chan error
+
 //go:linkname gracefulShutdownSignalNotify 
dubbo.apache.org/dubbo-go/v3/graceful_shutdown.signalNotify
 var gracefulShutdownSignalNotify func(chan<- os.Signal, ...os.Signal)
 
@@ -433,6 +436,7 @@ func resetGracefulShutdownStateForTest(t *testing.T) {
        gracefulShutdownStarted = atomic.Bool{}
        gracefulShutdownDone = make(chan struct{})
        gracefulShutdownResult = nil
+       gracefulShutdownSignalError = make(chan error, 1)
        gracefulShutdownSignalNotify = signal.Notify
 }
 
@@ -479,6 +483,49 @@ func TestServeContextReturnsAfterContextCancellation(t 
*testing.T) {
        }
 }
 
+func TestServeContextReturnsInternalShutdownError(t *testing.T) {
+       resetGracefulShutdownStateForTest(t)
+       t.Cleanup(func() {
+               resetGracefulShutdownStateForTest(t)
+       })
+       resetInternalProviderServicesForTest(t)
+       var registerCount atomic.Int32
+       registerCountingServeTestProtocols(t, nil, nil, &registerCount, nil, 
nil)
+
+       internalSignal := false
+       shutdownCfg := global.DefaultShutdownConfig()
+       shutdownCfg.InternalSignal = &internalSignal
+       shutdownCfg.ConsumerUpdateWaitTime = "0s"
+       shutdownCfg.StepTimeout = "0s"
+       shutdownCfg.NotifyTimeout = "10ms"
+       shutdownCfg.OfflineRequestWindowTimeout = "0s"
+
+       srv, err := NewServer(SetServerShutdown(shutdownCfg))
+       require.NoError(t, err)
+       require.NoError(t, srv.Register(&MockServerRPCService{}, nil))
+
+       serveDone := make(chan error, 1)
+       go func() {
+               serveDone <- srv.ServeContext(context.Background())
+       }()
+
+       require.Eventually(t, func() bool {
+               return registerCount.Load() == 1
+       }, time.Second, 10*time.Millisecond)
+
+       expectedErr := errors.New("graceful shutdown timed out")
+       gracefulShutdownSignalError <- expectedErr
+
+       select {
+       case err := <-serveDone:
+               require.ErrorIs(t, err, expectedErr)
+       case <-time.After(time.Second):
+               t.Fatal("ServeContext did not return the internal shutdown 
error")
+       }
+
+       require.NoError(t, graceful_shutdown.Shutdown(context.Background()))
+}
+
 func TestServeContextDoesNotStartWhenContextAlreadyCanceled(t *testing.T) {
        resetGracefulShutdownStateForTest(t)
        t.Cleanup(func() {

Reply via email to