From d4646b29815d4cd5e9c8fa529241f23dc4af1af9 Mon Sep 17 00:00:00 2001 From: Ryan Leung Date: Thu, 30 Jul 2026 12:06:57 +0800 Subject: [PATCH 1/3] server: return unavailable when PD is not serving Return gRPC Unavailable from service-discovery RPCs when the local PD server is not running. Preserve response-header errors for unrelated member-loading failures. Signed-off-by: Ryan Leung --- server/grpc_service.go | 35 +++++++++++++----------- server/grpc_service_test.go | 53 +++++++++++++++++++++++++++++++++++++ 2 files changed, 72 insertions(+), 16 deletions(-) diff --git a/server/grpc_service.go b/server/grpc_service.go index 87e1c3c97b9..0d32b7a09a9 100644 --- a/server/grpc_service.go +++ b/server/grpc_service.go @@ -249,9 +249,7 @@ func (s *GrpcServer) GetClusterInfo(context.Context, *pdpb.GetClusterInfoRequest // Here we purposely do not check the cluster ID because the client does not know the correct cluster ID // at startup and needs to get the cluster ID with the first request (i.e. GetMembers). if s.IsClosed() { - return &pdpb.GetClusterInfoResponse{ - Header: grpcutil.WrapErrorToHeader(pdpb.ErrorType_UNKNOWN, errs.ErrServerNotStarted.FastGenByArgs().Error()), - }, nil + return nil, errs.ErrNotStarted } var tsoServiceAddrs []string @@ -428,6 +426,11 @@ func (s *GrpcServer) getMinTSFromSingleServer( // GetMembers implements gRPC PDServer. func (s *GrpcServer) GetMembers(context.Context, *pdpb.GetMembersRequest) (*pdpb.GetMembersResponse, error) { + // Here we purposely do not check the cluster ID because the client does not know the correct cluster ID + // at startup and needs to get the cluster ID with the first request (i.e. GetMembers). + if s.IsClosed() { + return nil, errs.ErrNotStarted + } done, err := s.rateLimitCheck() if err != nil { return nil, err @@ -435,18 +438,9 @@ func (s *GrpcServer) GetMembers(context.Context, *pdpb.GetMembersRequest) (*pdpb if done != nil { defer done() } - // Here we purposely do not check the cluster ID because the client does not know the correct cluster ID - // at startup and needs to get the cluster ID with the first request (i.e. GetMembers). - if s.IsClosed() { - return &pdpb.GetMembersResponse{ - Header: grpcutil.WrapErrorToHeader(pdpb.ErrorType_UNKNOWN, errs.ErrServerNotStarted.FastGenByArgs().Error()), - }, nil - } members, err := s.Server.GetMembers() if err != nil { - return &pdpb.GetMembersResponse{ - Header: grpcutil.WrapErrorToHeader(pdpb.ErrorType_UNKNOWN, err.Error()), - }, nil + return getMembersErrorResult(err) } var etcdLeader, pdLeader *pdpb.Member @@ -456,9 +450,7 @@ func (s *GrpcServer) GetMembers(context.Context, *pdpb.GetMembersRequest) (*pdpb if (pdLeader == nil && leader != nil) || (etcdLeader == nil && leaderID != 0) { members, err = s.ReloadMembers() if err != nil { - return &pdpb.GetMembersResponse{ - Header: grpcutil.WrapErrorToHeader(pdpb.ErrorType_UNKNOWN, err.Error()), - }, nil + return getMembersErrorResult(err) } leaderID = s.member.GetEtcdLeader() leader = s.member.GetLeader() @@ -473,6 +465,17 @@ func (s *GrpcServer) GetMembers(context.Context, *pdpb.GetMembersRequest) (*pdpb }, nil } +func getMembersErrorResult(err error) (*pdpb.GetMembersResponse, error) { + if errors.ErrorEqual(err, errs.ErrServerNotStarted.FastGenByArgs()) { + return nil, errs.ErrNotStarted + } + // Keep other server-side failures in the response header so clients don't + // mistake them for service or transport availability failures. + return &pdpb.GetMembersResponse{ + Header: grpcutil.WrapErrorToHeader(pdpb.ErrorType_UNKNOWN, err.Error()), + }, nil +} + func findLeadersInMembers(members []*pdpb.Member, etcdLeaderID uint64, leader *pdpb.Member) (etcdLeader, pdLeader *pdpb.Member) { for _, m := range members { if m.GetMemberId() == etcdLeaderID { diff --git a/server/grpc_service_test.go b/server/grpc_service_test.go index 040d944fd54..748a0371a67 100644 --- a/server/grpc_service_test.go +++ b/server/grpc_service_test.go @@ -15,15 +15,24 @@ package server import ( + "context" + "net" "testing" "github.com/stretchr/testify/require" "go.uber.org/goleak" + "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/credentials/insecure" + "google.golang.org/grpc/status" + "google.golang.org/grpc/test/bufconn" + "github.com/pingcap/errors" "github.com/pingcap/kvproto/pkg/metapb" "github.com/pingcap/kvproto/pkg/pdpb" "github.com/pingcap/kvproto/pkg/schedulingpb" + "github.com/tikv/pd/pkg/errs" "github.com/tikv/pd/pkg/utils/testutil" ) @@ -94,3 +103,47 @@ func TestConvertSchedulingHeaderPreservesError(t *testing.T) { }) } } + +func TestServiceDiscoveryRPCsReturnUnavailableWhenServerIsNotRunning(t *testing.T) { + grpcServer := &GrpcServer{Server: &Server{}} + listener := bufconn.Listen(1024 * 1024) + transport := grpc.NewServer() + pdpb.RegisterPDServer(transport, grpcServer) + serveErr := make(chan error, 1) + go func() { + serveErr <- transport.Serve(listener) + }() + t.Cleanup(func() { + transport.Stop() + require.NoError(t, <-serveErr) + }) + conn, err := grpc.NewClient( + "passthrough:///bufnet", + grpc.WithContextDialer(func(context.Context, string) (net.Conn, error) { return listener.Dial() }), + grpc.WithTransportCredentials(insecure.NewCredentials()), + ) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, conn.Close()) }) + client := pdpb.NewPDClient(conn) + + members, err := client.GetMembers(context.Background(), &pdpb.GetMembersRequest{}) + require.Nil(t, members) + require.Equal(t, codes.Unavailable, status.Code(err)) + + clusterInfo, err := client.GetClusterInfo(context.Background(), &pdpb.GetClusterInfoRequest{}) + require.Nil(t, clusterInfo) + require.Equal(t, codes.Unavailable, status.Code(err)) +} + +func TestGetMembersErrorResult(t *testing.T) { + notStarted := errors.WithStack(errs.ErrServerNotStarted.FastGenByArgs()) + response, err := getMembersErrorResult(notStarted) + require.Nil(t, response) + require.Equal(t, codes.Unavailable, status.Code(err)) + + internalErr := errors.New("failed to load members") + response, err = getMembersErrorResult(internalErr) + require.NoError(t, err) + require.Equal(t, pdpb.ErrorType_UNKNOWN, response.GetHeader().GetError().GetType()) + require.Equal(t, internalErr.Error(), response.GetHeader().GetError().GetMessage()) +} From 2e5b0f724f94c78fe4313060f213145f06fe9db5 Mon Sep 17 00:00:00 2001 From: Ryan Leung Date: Thu, 30 Jul 2026 15:15:50 +0800 Subject: [PATCH 2/3] mcs/tso: return unavailable when service is not running Signed-off-by: Ryan Leung --- pkg/mcs/tso/server/grpc_service.go | 8 +++- pkg/mcs/tso/server/grpc_service_test.go | 61 +++++++++++++++++++++++++ 2 files changed, 68 insertions(+), 1 deletion(-) create mode 100644 pkg/mcs/tso/server/grpc_service_test.go diff --git a/pkg/mcs/tso/server/grpc_service.go b/pkg/mcs/tso/server/grpc_service.go index 56d957f7b04..546990e6217 100644 --- a/pkg/mcs/tso/server/grpc_service.go +++ b/pkg/mcs/tso/server/grpc_service.go @@ -136,6 +136,9 @@ func (s *Service) Tso(stream tsopb.TSO_TsoServer) error { func (s *Service) FindGroupByKeyspaceID( _ context.Context, request *tsopb.FindGroupByKeyspaceIDRequest, ) (*tsopb.FindGroupByKeyspaceIDResponse, error) { + if s.IsClosed() { + return nil, errs.ErrNotStarted + } respKeyspaceGroup := request.GetHeader().GetKeyspaceGroupId() if errorType, err := s.validRequest(request.GetHeader()); err != nil { return &tsopb.FindGroupByKeyspaceIDResponse{ @@ -198,6 +201,9 @@ func (s *Service) FindGroupByKeyspaceID( func (s *Service) GetMinTS( _ context.Context, request *tsopb.GetMinTSRequest, ) (*tsopb.GetMinTSResponse, error) { + if s.IsClosed() { + return nil, errs.ErrNotStarted + } respKeyspaceGroup := request.GetHeader().GetKeyspaceGroupId() if errorType, err := s.validRequest(request.GetHeader()); err != nil { return &tsopb.GetMinTSResponse{ @@ -225,7 +231,7 @@ func (s *Service) GetMinTS( } func (s *Service) validRequest(header *tsopb.RequestHeader) (tsopb.ErrorType, error) { - if s.IsClosed() || s.keyspaceGroupManager == nil { + if s.keyspaceGroupManager == nil { return tsopb.ErrorType_NOT_BOOTSTRAPPED, errs.ErrNotStarted } if header == nil || header.GetClusterId() != keypath.ClusterID() { diff --git a/pkg/mcs/tso/server/grpc_service_test.go b/pkg/mcs/tso/server/grpc_service_test.go new file mode 100644 index 00000000000..6425f891041 --- /dev/null +++ b/pkg/mcs/tso/server/grpc_service_test.go @@ -0,0 +1,61 @@ +// Copyright 2026 TiKV Project Authors. +// +// Licensed 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 server + +import ( + "context" + "net" + "testing" + + "github.com/stretchr/testify/require" + "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/credentials/insecure" + "google.golang.org/grpc/status" + "google.golang.org/grpc/test/bufconn" + + "github.com/pingcap/kvproto/pkg/tsopb" +) + +func TestUnaryRPCsReturnUnavailableWhenServerIsNotRunning(t *testing.T) { + service := &Service{Server: &Server{}} + listener := bufconn.Listen(1024 * 1024) + transport := grpc.NewServer() + tsopb.RegisterTSOServer(transport, service) + serveErr := make(chan error, 1) + go func() { + serveErr <- transport.Serve(listener) + }() + t.Cleanup(func() { + transport.Stop() + require.NoError(t, <-serveErr) + }) + conn, err := grpc.NewClient( + "passthrough:///bufnet", + grpc.WithContextDialer(func(context.Context, string) (net.Conn, error) { return listener.Dial() }), + grpc.WithTransportCredentials(insecure.NewCredentials()), + ) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, conn.Close()) }) + client := tsopb.NewTSOClient(conn) + + group, err := client.FindGroupByKeyspaceID(context.Background(), &tsopb.FindGroupByKeyspaceIDRequest{}) + require.Nil(t, group) + require.Equal(t, codes.Unavailable, status.Code(err)) + + minTS, err := client.GetMinTS(context.Background(), &tsopb.GetMinTSRequest{}) + require.Nil(t, minTS) + require.Equal(t, codes.Unavailable, status.Code(err)) +} From 46eba19c802be92ef504a647a8355494de8bf960 Mon Sep 17 00:00:00 2001 From: Ryan Leung Date: Thu, 30 Jul 2026 17:50:19 +0800 Subject: [PATCH 3/3] server: keep GetMembers rate limiting first Signed-off-by: Ryan Leung --- server/grpc_service.go | 10 +++++----- server/grpc_service_test.go | 5 ++++- 2 files changed, 9 insertions(+), 6 deletions(-) diff --git a/server/grpc_service.go b/server/grpc_service.go index 0d32b7a09a9..66f673d503d 100644 --- a/server/grpc_service.go +++ b/server/grpc_service.go @@ -426,11 +426,6 @@ func (s *GrpcServer) getMinTSFromSingleServer( // GetMembers implements gRPC PDServer. func (s *GrpcServer) GetMembers(context.Context, *pdpb.GetMembersRequest) (*pdpb.GetMembersResponse, error) { - // Here we purposely do not check the cluster ID because the client does not know the correct cluster ID - // at startup and needs to get the cluster ID with the first request (i.e. GetMembers). - if s.IsClosed() { - return nil, errs.ErrNotStarted - } done, err := s.rateLimitCheck() if err != nil { return nil, err @@ -438,6 +433,11 @@ func (s *GrpcServer) GetMembers(context.Context, *pdpb.GetMembersRequest) (*pdpb if done != nil { defer done() } + // Here we purposely do not check the cluster ID because the client does not know the correct cluster ID + // at startup and needs to get the cluster ID with the first request (i.e. GetMembers). + if s.IsClosed() { + return nil, errs.ErrNotStarted + } members, err := s.Server.GetMembers() if err != nil { return getMembersErrorResult(err) diff --git a/server/grpc_service_test.go b/server/grpc_service_test.go index 748a0371a67..909437e314a 100644 --- a/server/grpc_service_test.go +++ b/server/grpc_service_test.go @@ -34,6 +34,7 @@ import ( "github.com/tikv/pd/pkg/errs" "github.com/tikv/pd/pkg/utils/testutil" + "github.com/tikv/pd/server/config" ) func TestMain(m *testing.M) { @@ -105,7 +106,9 @@ func TestConvertSchedulingHeaderPreservesError(t *testing.T) { } func TestServiceDiscoveryRPCsReturnUnavailableWhenServerIsNotRunning(t *testing.T) { - grpcServer := &GrpcServer{Server: &Server{}} + grpcServer := &GrpcServer{Server: &Server{ + serviceMiddlewarePersistOptions: config.NewServiceMiddlewarePersistOptions(&config.ServiceMiddlewareConfig{}), + }} listener := bufconn.Listen(1024 * 1024) transport := grpc.NewServer() pdpb.RegisterPDServer(transport, grpcServer)