From 993b7a20097e07bcb4cf274fd2f46e320758749d Mon Sep 17 00:00:00 2001 From: Ryan Leung Date: Thu, 30 Jul 2026 18:24:21 +0800 Subject: [PATCH 1/8] server: validate forwarded PD hosts before dialing Signed-off-by: Ryan Leung --- server/forward.go | 35 +++++++++++++++++++++++ server/grpc_service.go | 8 +++--- server/resource_group_proxy_service.go | 2 +- tests/server/tso/tso_proxy_test.go | 39 ++++++++++++++++++++++++++ 4 files changed, 79 insertions(+), 5 deletions(-) diff --git a/server/forward.go b/server/forward.go index a5d97bfe81e..a20188b824c 100644 --- a/server/forward.go +++ b/server/forward.go @@ -21,6 +21,8 @@ import ( "go.uber.org/zap" "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" "github.com/pingcap/errors" "github.com/pingcap/failpoint" @@ -450,6 +452,39 @@ func (s *GrpcServer) getDelegateClient(ctx context.Context, forwardedHost string return conn.(*grpc.ClientConn), nil } +// getPDForwardedDelegateClient returns a delegate client for a PD address received +// from the forwarding metadata. Unlike addresses discovered by PD itself, the +// metadata is controlled by the caller, so it must be restricted to the advertised +// client URLs of the current PD members before dialing. +func (s *GrpcServer) getPDForwardedDelegateClient(ctx context.Context, forwardedHost string) (*grpc.ClientConn, error) { + if memberHasClientURL(s.GetLeader(), forwardedHost) { + return s.getDelegateClient(ctx, forwardedHost) + } + + members, err := s.Server.GetMembers() + if err != nil { + return nil, status.Errorf(codes.Unavailable, "failed to get PD members: %v", err) + } + for _, member := range members { + if memberHasClientURL(member, forwardedHost) { + return s.getDelegateClient(ctx, forwardedHost) + } + } + return nil, status.Errorf(codes.InvalidArgument, "forwarded host %q is not a client URL of any PD member", forwardedHost) +} + +func memberHasClientURL(member *pdpb.Member, addr string) bool { + if member == nil { + return false + } + for _, clientURL := range member.GetClientUrls() { + if clientURL == addr { + return true + } + } + return false +} + func (s *GrpcServer) closeDelegateClient(forwardedHost string) { client, ok := s.clientConns.LoadAndDelete(forwardedHost) if !ok { diff --git a/server/grpc_service.go b/server/grpc_service.go index 87e1c3c97b9..b7370bb3d5f 100644 --- a/server/grpc_service.go +++ b/server/grpc_service.go @@ -231,7 +231,7 @@ func (s *GrpcServer) unaryFollowerMiddleware(ctx context.Context, req request, f }) forwardedHost := grpcutil.GetForwardedHost(ctx) if !s.isLocalRequest(forwardedHost) { - client, err := s.getDelegateClient(ctx, forwardedHost) + client, err := s.getPDForwardedDelegateClient(ctx, forwardedHost) if err != nil { return nil, err } @@ -565,7 +565,7 @@ func (s *GrpcServer) Tso(stream pdpb.PD_TsoServer) error { forwardedHost := grpcutil.GetForwardedHost(stream.Context()) if !s.isLocalRequest(forwardedHost) { - clientConn, err := s.getDelegateClient(s.ctx, forwardedHost) + clientConn, err := s.getPDForwardedDelegateClient(s.ctx, forwardedHost) if err != nil { return errors.WithStack(err) } @@ -1114,7 +1114,7 @@ func (s *GrpcServer) ReportBuckets(stream pdpb.PD_ReportBucketsServer) error { if cancel != nil { cancel() } - client, err := s.getDelegateClient(s.ctx, forwardedHost) + client, err := s.getPDForwardedDelegateClient(s.ctx, forwardedHost) if err != nil { return err } @@ -1292,7 +1292,7 @@ func (s *GrpcServer) RegionHeartbeat(stream pdpb.PD_RegionHeartbeatServer) error if cancel != nil { cancel() } - client, err := s.getDelegateClient(s.ctx, forwardedHost) + client, err := s.getPDForwardedDelegateClient(s.ctx, forwardedHost) if err != nil { return err } diff --git a/server/resource_group_proxy_service.go b/server/resource_group_proxy_service.go index 72f1f249aea..36c581168dc 100644 --- a/server/resource_group_proxy_service.go +++ b/server/resource_group_proxy_service.go @@ -72,7 +72,7 @@ func (s *resourceGroupProxyServer) getPDMetadataWriteDelegateClient(ctx context. if s.isLocalRequest(forwardedHost) { return nil, "", nil } - client, err := s.getDelegateClient(ctx, forwardedHost) + client, err := s.getPDForwardedDelegateClient(ctx, forwardedHost) if err != nil { return nil, "", err } diff --git a/tests/server/tso/tso_proxy_test.go b/tests/server/tso/tso_proxy_test.go index 5ad69d25fcc..bfe61d93b5a 100644 --- a/tests/server/tso/tso_proxy_test.go +++ b/tests/server/tso/tso_proxy_test.go @@ -17,12 +17,15 @@ package tso_test import ( "context" "io" + "net" "testing" "time" "github.com/stretchr/testify/require" "github.com/stretchr/testify/suite" "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" "github.com/pingcap/failpoint" "github.com/pingcap/kvproto/pkg/pdpb" @@ -138,6 +141,42 @@ func (s *tsoProxyTestSuite) verifyProxyIsHealthyWith(client pdpb.PD_TsoClient) { re.GreaterOrEqual(uint32(timestamp.GetLogical()), s.defaultReq.GetCount()) } +func (s *tsoProxyTestSuite) TestRejectUnknownForwardedHost() { + re := s.Require() + listener, err := net.Listen("tcp", "127.0.0.1:0") + re.NoError(err) + accepted := make(chan bool, 1) + go func() { + conn, err := listener.Accept() + if err != nil { + accepted <- false + return + } + conn.Close() + accepted <- true + }() + defer func() { + re.NoError(listener.Close()) + re.False(<-accepted) + }() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + ctx = grpcutil.BuildForwardContext(ctx, "http://"+listener.Addr().String()) + + _, err = s.pdClient.GetAllStores(ctx, &pdpb.GetAllStoresRequest{Header: s.defaultReq.GetHeader()}) + re.Error(err) + re.Equal(codes.InvalidArgument, status.Code(err)) + + client, err := s.pdClient.Tso(ctx) + re.NoError(err) + defer client.CloseSend() + re.NoError(client.Send(s.defaultReq)) + _, err = client.Recv() + re.Error(err) + re.Equal(codes.InvalidArgument, status.Code(err)) +} + func (s *tsoProxyTestSuite) assertReceiveError(re *require.Assertions, errStr string) { re.NoError(s.proxyClient.Send(s.defaultReq)) _, err := s.proxyClient.Recv() From 152970b4770a12d10e0414c11be37ee4082f3917 Mon Sep 17 00:00:00 2001 From: Ryan Leung Date: Thu, 30 Jul 2026 18:45:35 +0800 Subject: [PATCH 2/8] server: restrict forwarded targets to current PD leader Signed-off-by: Ryan Leung --- server/forward.go | 35 ++++++++---------------------- server/grpc_service.go | 19 +++++++++------- tests/server/tso/tso_proxy_test.go | 11 ++++++++++ 3 files changed, 31 insertions(+), 34 deletions(-) diff --git a/server/forward.go b/server/forward.go index a20188b824c..da647f4966b 100644 --- a/server/forward.go +++ b/server/forward.go @@ -452,37 +452,20 @@ func (s *GrpcServer) getDelegateClient(ctx context.Context, forwardedHost string return conn.(*grpc.ClientConn), nil } -// getPDForwardedDelegateClient returns a delegate client for a PD address received -// from the forwarding metadata. Unlike addresses discovered by PD itself, the -// metadata is controlled by the caller, so it must be restricted to the advertised -// client URLs of the current PD members before dialing. +// getPDForwardedDelegateClient returns a delegate client for the PD leader address +// received from the forwarding metadata. The metadata is controlled by the caller, +// so it must match an advertised client URL of the current PD leader before dialing. func (s *GrpcServer) getPDForwardedDelegateClient(ctx context.Context, forwardedHost string) (*grpc.ClientConn, error) { - if memberHasClientURL(s.GetLeader(), forwardedHost) { - return s.getDelegateClient(ctx, forwardedHost) + leader := s.GetLeader() + if leader == nil || len(leader.GetClientUrls()) == 0 { + return nil, status.Error(codes.Unavailable, "PD leader is not available") } - - members, err := s.Server.GetMembers() - if err != nil { - return nil, status.Errorf(codes.Unavailable, "failed to get PD members: %v", err) - } - for _, member := range members { - if memberHasClientURL(member, forwardedHost) { + for _, clientURL := range leader.GetClientUrls() { + if clientURL == forwardedHost { return s.getDelegateClient(ctx, forwardedHost) } } - return nil, status.Errorf(codes.InvalidArgument, "forwarded host %q is not a client URL of any PD member", forwardedHost) -} - -func memberHasClientURL(member *pdpb.Member, addr string) bool { - if member == nil { - return false - } - for _, clientURL := range member.GetClientUrls() { - if clientURL == addr { - return true - } - } - return false + return nil, status.Errorf(codes.InvalidArgument, "forwarded host %q is not a client URL of the PD leader", forwardedHost) } func (s *GrpcServer) closeDelegateClient(forwardedHost string) { diff --git a/server/grpc_service.go b/server/grpc_service.go index b7370bb3d5f..bcff2d15855 100644 --- a/server/grpc_service.go +++ b/server/grpc_service.go @@ -504,9 +504,11 @@ func (s *GrpcServer) Tso(stream pdpb.PD_TsoServer) error { var ( // The following are tso forward stream related variables. - tsoRequestProxyCtx context.Context - forwarder = newTSOForwarder(stream) - tsoStreamErr error + tsoRequestProxyCtx context.Context + forwardedHost = grpcutil.GetForwardedHost(stream.Context()) + forwardedClientConn *grpc.ClientConn + forwarder = newTSOForwarder(stream) + tsoStreamErr error ) defer func() { @@ -563,14 +565,15 @@ func (s *GrpcServer) Tso(stream pdpb.PD_TsoServer) error { return errs.ErrNotStarted } - forwardedHost := grpcutil.GetForwardedHost(stream.Context()) if !s.isLocalRequest(forwardedHost) { - clientConn, err := s.getPDForwardedDelegateClient(s.ctx, forwardedHost) - if err != nil { - return errors.WithStack(err) + if forwardedClientConn == nil { + forwardedClientConn, err = s.getPDForwardedDelegateClient(s.ctx, forwardedHost) + if err != nil { + return errors.WithStack(err) + } } - tsoRequest := tsoutil.NewPDProtoRequest(forwardedHost, clientConn, request, stream) + tsoRequest := tsoutil.NewPDProtoRequest(forwardedHost, forwardedClientConn, request, stream) // don't pass a stream context here as dispatcher serves multiple streams tsoRequestProxyCtx = s.tsoDispatcher.DispatchRequest(s.ctx, tsoRequest, s.pdProtoFactory, s.tsoPrimaryWatcher) continue diff --git a/tests/server/tso/tso_proxy_test.go b/tests/server/tso/tso_proxy_test.go index bfe61d93b5a..e06eaed9e2f 100644 --- a/tests/server/tso/tso_proxy_test.go +++ b/tests/server/tso/tso_proxy_test.go @@ -141,6 +141,17 @@ func (s *tsoProxyTestSuite) verifyProxyIsHealthyWith(client pdpb.PD_TsoClient) { re.GreaterOrEqual(uint32(timestamp.GetLogical()), s.defaultReq.GetCount()) } +func (s *tsoProxyTestSuite) TestRejectFollowerForwardedHost() { + re := s.Require() + client, conn := testutil.MustNewGrpcClient(re, s.leader.GetAddr()) + defer conn.Close() + + ctx := grpcutil.BuildForwardContext(context.Background(), s.follower.GetAddr()) + _, err := client.GetAllStores(ctx, &pdpb.GetAllStoresRequest{Header: s.defaultReq.GetHeader()}) + re.Error(err) + re.Equal(codes.InvalidArgument, status.Code(err)) +} + func (s *tsoProxyTestSuite) TestRejectUnknownForwardedHost() { re := s.Require() listener, err := net.Listen("tcp", "127.0.0.1:0") From 9eb68aaf0374432dc44e23c67986729b0f0b80a7 Mon Sep 17 00:00:00 2001 From: Ryan Leung Date: Fri, 31 Jul 2026 10:53:44 +0800 Subject: [PATCH 3/8] tests: check TSO stream close error Signed-off-by: Ryan Leung --- tests/server/tso/tso_proxy_test.go | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/tests/server/tso/tso_proxy_test.go b/tests/server/tso/tso_proxy_test.go index e06eaed9e2f..233517f34dd 100644 --- a/tests/server/tso/tso_proxy_test.go +++ b/tests/server/tso/tso_proxy_test.go @@ -181,7 +181,12 @@ func (s *tsoProxyTestSuite) TestRejectUnknownForwardedHost() { client, err := s.pdClient.Tso(ctx) re.NoError(err) - defer client.CloseSend() + defer func() { + err := client.CloseSend() + if err != nil && err != io.EOF { + re.NoError(err) + } + }() re.NoError(client.Send(s.defaultReq)) _, err = client.Recv() re.Error(err) From dd9fc10b71c6c5961df16924cf061440d239dbb0 Mon Sep 17 00:00:00 2001 From: Ryan Leung Date: Fri, 31 Jul 2026 11:20:18 +0800 Subject: [PATCH 4/8] tests: cover follower TSO forwarding rejection Signed-off-by: Ryan Leung --- tests/server/tso/tso_proxy_test.go | 17 ++++++++++++++++- 1 file changed, 16 insertions(+), 1 deletion(-) diff --git a/tests/server/tso/tso_proxy_test.go b/tests/server/tso/tso_proxy_test.go index 233517f34dd..41508e6efba 100644 --- a/tests/server/tso/tso_proxy_test.go +++ b/tests/server/tso/tso_proxy_test.go @@ -146,10 +146,25 @@ func (s *tsoProxyTestSuite) TestRejectFollowerForwardedHost() { client, conn := testutil.MustNewGrpcClient(re, s.leader.GetAddr()) defer conn.Close() - ctx := grpcutil.BuildForwardContext(context.Background(), s.follower.GetAddr()) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + ctx = grpcutil.BuildForwardContext(ctx, s.follower.GetAddr()) _, err := client.GetAllStores(ctx, &pdpb.GetAllStoresRequest{Header: s.defaultReq.GetHeader()}) re.Error(err) re.Equal(codes.InvalidArgument, status.Code(err)) + + tsoClient, err := client.Tso(ctx) + re.NoError(err) + defer func() { + err := tsoClient.CloseSend() + if err != nil && err != io.EOF { + re.NoError(err) + } + }() + re.NoError(tsoClient.Send(s.defaultReq)) + _, err = tsoClient.Recv() + re.Error(err) + re.Equal(codes.InvalidArgument, status.Code(err)) } func (s *tsoProxyTestSuite) TestRejectUnknownForwardedHost() { From 0a49a48f0e5fe70c33b798a12f63a893e1154dec Mon Sep 17 00:00:00 2001 From: Ryan Leung Date: Fri, 31 Jul 2026 13:49:50 +0800 Subject: [PATCH 5/8] server: preserve internal metadata write forwarding Signed-off-by: Ryan Leung --- server/resource_group_proxy_service.go | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/server/resource_group_proxy_service.go b/server/resource_group_proxy_service.go index 36c581168dc..65bfc96f037 100644 --- a/server/resource_group_proxy_service.go +++ b/server/resource_group_proxy_service.go @@ -62,6 +62,7 @@ func (s *resourceGroupProxyServer) closeClient(ctx context.Context) { func (s *resourceGroupProxyServer) getPDMetadataWriteDelegateClient(ctx context.Context) (resource_manager.ResourceManagerClient, string, error) { forwardedHost := grpcutil.GetForwardedHost(ctx) + forwardedHostFromMetadata := forwardedHost != "" if forwardedHost == "" { leader := s.GetLeader() if leader == nil || len(leader.GetClientUrls()) == 0 { @@ -72,6 +73,14 @@ func (s *resourceGroupProxyServer) getPDMetadataWriteDelegateClient(ctx context. if s.isLocalRequest(forwardedHost) { return nil, "", nil } + if !forwardedHostFromMetadata { + // Keep PD-discovered leader targets on the original trusted delegate path. + client, err := s.getDelegateClient(ctx, forwardedHost) + if err != nil { + return nil, "", err + } + return resource_manager.NewResourceManagerClient(client), forwardedHost, nil + } client, err := s.getPDForwardedDelegateClient(ctx, forwardedHost) if err != nil { return nil, "", err From 4e1a1ff882c81bf0b88ef15ead53a0f21e79c5a1 Mon Sep 17 00:00:00 2001 From: Ryan Leung Date: Fri, 31 Jul 2026 14:41:42 +0800 Subject: [PATCH 6/8] server: validate forwarded metadata before local handling Signed-off-by: Ryan Leung --- server/forward.go | 14 +++--- server/grpc_service.go | 50 +++++++++++++++---- server/resource_group_proxy_service.go | 17 +++---- .../mcs/resourcemanager/redirector_test.go | 18 +++++-- tests/server/tso/tso_proxy_test.go | 8 ++- 5 files changed, 74 insertions(+), 33 deletions(-) diff --git a/server/forward.go b/server/forward.go index da647f4966b..525a6e3c9aa 100644 --- a/server/forward.go +++ b/server/forward.go @@ -452,20 +452,20 @@ func (s *GrpcServer) getDelegateClient(ctx context.Context, forwardedHost string return conn.(*grpc.ClientConn), nil } -// getPDForwardedDelegateClient returns a delegate client for the PD leader address -// received from the forwarding metadata. The metadata is controlled by the caller, -// so it must match an advertised client URL of the current PD leader before dialing. -func (s *GrpcServer) getPDForwardedDelegateClient(ctx context.Context, forwardedHost string) (*grpc.ClientConn, error) { +// validatePDForwardedHost checks that a PD address received from forwarding +// metadata belongs to the current PD leader. The metadata is controlled by the +// caller, so it must be validated before both local handling and dialing. +func (s *GrpcServer) validatePDForwardedHost(forwardedHost string) error { leader := s.GetLeader() if leader == nil || len(leader.GetClientUrls()) == 0 { - return nil, status.Error(codes.Unavailable, "PD leader is not available") + return status.Error(codes.Unavailable, "PD leader is not available") } for _, clientURL := range leader.GetClientUrls() { if clientURL == forwardedHost { - return s.getDelegateClient(ctx, forwardedHost) + return nil } } - return nil, status.Errorf(codes.InvalidArgument, "forwarded host %q is not a client URL of the PD leader", forwardedHost) + return status.Errorf(codes.InvalidArgument, "forwarded host %q is not a client URL of the PD leader", forwardedHost) } func (s *GrpcServer) closeDelegateClient(forwardedHost string) { diff --git a/server/grpc_service.go b/server/grpc_service.go index bcff2d15855..b68106a39ab 100644 --- a/server/grpc_service.go +++ b/server/grpc_service.go @@ -230,8 +230,13 @@ func (s *GrpcServer) unaryFollowerMiddleware(ctx context.Context, req request, f time.Sleep(5 * time.Second) }) forwardedHost := grpcutil.GetForwardedHost(ctx) + if forwardedHost != "" { + if err := s.validatePDForwardedHost(forwardedHost); err != nil { + return nil, err + } + } if !s.isLocalRequest(forwardedHost) { - client, err := s.getPDForwardedDelegateClient(ctx, forwardedHost) + client, err := s.getDelegateClient(ctx, forwardedHost) if err != nil { return nil, err } @@ -504,11 +509,12 @@ func (s *GrpcServer) Tso(stream pdpb.PD_TsoServer) error { var ( // The following are tso forward stream related variables. - tsoRequestProxyCtx context.Context - forwardedHost = grpcutil.GetForwardedHost(stream.Context()) - forwardedClientConn *grpc.ClientConn - forwarder = newTSOForwarder(stream) - tsoStreamErr error + tsoRequestProxyCtx context.Context + forwardedHost = grpcutil.GetForwardedHost(stream.Context()) + forwardedClientConn *grpc.ClientConn + forwarder = newTSOForwarder(stream) + tsoStreamErr error + forwardedHostChecked = forwardedHost == "" ) defer func() { @@ -565,9 +571,15 @@ func (s *GrpcServer) Tso(stream pdpb.PD_TsoServer) error { return errs.ErrNotStarted } + if !forwardedHostChecked { + if err := s.validatePDForwardedHost(forwardedHost); err != nil { + return errors.WithStack(err) + } + forwardedHostChecked = true + } if !s.isLocalRequest(forwardedHost) { if forwardedClientConn == nil { - forwardedClientConn, err = s.getPDForwardedDelegateClient(s.ctx, forwardedHost) + forwardedClientConn, err = s.getDelegateClient(s.ctx, forwardedHost) if err != nil { return errors.WithStack(err) } @@ -1083,6 +1095,8 @@ func (s *GrpcServer) ReportBuckets(stream pdpb.PD_ReportBucketsServer) error { forwardErrCh chan error forwardSchedulingStream schedulingpb.Scheduling_RegionBucketsClient lastForwardedSchedulingHost string + metadataForwardedHost = grpcutil.GetForwardedHost(stream.Context()) + forwardedHostChecked = metadataForwardedHost == "" ) defer func() { if cancel != nil { @@ -1108,7 +1122,13 @@ func (s *GrpcServer) ReportBuckets(stream pdpb.PD_ReportBucketsServer) error { if err != nil { return errors.WithStack(err) } - forwardedHost := grpcutil.GetForwardedHost(stream.Context()) + if !forwardedHostChecked { + if err := s.validatePDForwardedHost(metadataForwardedHost); err != nil { + return err + } + forwardedHostChecked = true + } + forwardedHost := metadataForwardedHost failpoint.Inject("grpcClientClosed", func() { forwardedHost = s.GetMember().Member().GetClientUrls()[0] }) @@ -1117,7 +1137,7 @@ func (s *GrpcServer) ReportBuckets(stream pdpb.PD_ReportBucketsServer) error { if cancel != nil { cancel() } - client, err := s.getPDForwardedDelegateClient(s.ctx, forwardedHost) + client, err := s.getDelegateClient(s.ctx, forwardedHost) if err != nil { return err } @@ -1264,6 +1284,8 @@ func (s *GrpcServer) RegionHeartbeat(stream pdpb.PD_RegionHeartbeatServer) error forwardErrCh chan error forwardSchedulingStream schedulingpb.Scheduling_RegionHeartbeatClient lastForwardedSchedulingHost string + metadataForwardedHost = grpcutil.GetForwardedHost(stream.Context()) + forwardedHostChecked = metadataForwardedHost == "" ) defer func() { // cancel the forward stream @@ -1286,7 +1308,13 @@ func (s *GrpcServer) RegionHeartbeat(stream pdpb.PD_RegionHeartbeatServer) error if err != nil { return errors.WithStack(err) } - forwardedHost := grpcutil.GetForwardedHost(stream.Context()) + if !forwardedHostChecked { + if err := s.validatePDForwardedHost(metadataForwardedHost); err != nil { + return err + } + forwardedHostChecked = true + } + forwardedHost := metadataForwardedHost failpoint.Inject("grpcClientClosed", func() { forwardedHost = s.GetMember().Member().GetClientUrls()[0] }) @@ -1295,7 +1323,7 @@ func (s *GrpcServer) RegionHeartbeat(stream pdpb.PD_RegionHeartbeatServer) error if cancel != nil { cancel() } - client, err := s.getPDForwardedDelegateClient(s.ctx, forwardedHost) + client, err := s.getDelegateClient(s.ctx, forwardedHost) if err != nil { return err } diff --git a/server/resource_group_proxy_service.go b/server/resource_group_proxy_service.go index 65bfc96f037..d44be436f4e 100644 --- a/server/resource_group_proxy_service.go +++ b/server/resource_group_proxy_service.go @@ -62,8 +62,11 @@ func (s *resourceGroupProxyServer) closeClient(ctx context.Context) { func (s *resourceGroupProxyServer) getPDMetadataWriteDelegateClient(ctx context.Context) (resource_manager.ResourceManagerClient, string, error) { forwardedHost := grpcutil.GetForwardedHost(ctx) - forwardedHostFromMetadata := forwardedHost != "" - if forwardedHost == "" { + if forwardedHost != "" { + if err := s.validatePDForwardedHost(forwardedHost); err != nil { + return nil, "", err + } + } else { leader := s.GetLeader() if leader == nil || len(leader.GetClientUrls()) == 0 { return nil, "", status.Error(codes.Unavailable, "pd leader is not available") @@ -73,15 +76,7 @@ func (s *resourceGroupProxyServer) getPDMetadataWriteDelegateClient(ctx context. if s.isLocalRequest(forwardedHost) { return nil, "", nil } - if !forwardedHostFromMetadata { - // Keep PD-discovered leader targets on the original trusted delegate path. - client, err := s.getDelegateClient(ctx, forwardedHost) - if err != nil { - return nil, "", err - } - return resource_manager.NewResourceManagerClient(client), forwardedHost, nil - } - client, err := s.getPDForwardedDelegateClient(ctx, forwardedHost) + client, err := s.getDelegateClient(ctx, forwardedHost) if err != nil { return nil, "", err } diff --git a/tests/integrations/mcs/resourcemanager/redirector_test.go b/tests/integrations/mcs/resourcemanager/redirector_test.go index 09f4f046b6e..895192f87d0 100644 --- a/tests/integrations/mcs/resourcemanager/redirector_test.go +++ b/tests/integrations/mcs/resourcemanager/redirector_test.go @@ -347,14 +347,26 @@ func (suite *resourceManagerRedirectorTestSuite) TestGRPCMetadataWritesForwardFr group.Priority = 11 group.RUSettings.RU.Settings.FillRate = 960 group.RUSettings.RU.Settings.BurstLimit = 1024 - modifyResp, err := followerClient.ModifyResourceGroup(ctx, &rmpb.PutResourceGroupRequest{Group: group}) - re.NoError(err) - re.Equal("Success!", modifyResp.GetBody()) + forwardedCtx := grpcutil.BuildForwardContext(ctx, suite.pdFollower.GetAddr()) + _, err = followerClient.ModifyResourceGroup(forwardedCtx, &rmpb.PutResourceGroupRequest{Group: group}) + re.Error(err) + re.Equal(codes.InvalidArgument, status.Code(err)) getReq := &rmpb.GetResourceGroupRequest{ ResourceGroupName: groupName, KeyspaceId: &rmpb.KeyspaceIDValue{Value: suite.keyspaceID}, } + unmodifiedResp, err := leaderClient.GetResourceGroup(ctx, getReq) + re.NoError(err) + re.NotNil(unmodifiedResp.GetGroup()) + re.Equal(uint32(5), unmodifiedResp.GetGroup().GetPriority()) + re.Equal(uint64(320), unmodifiedResp.GetGroup().GetRUSettings().GetRU().GetSettings().GetFillRate()) + re.Equal(int64(480), unmodifiedResp.GetGroup().GetRUSettings().GetRU().GetSettings().GetBurstLimit()) + + modifyResp, err := followerClient.ModifyResourceGroup(ctx, &rmpb.PutResourceGroupRequest{Group: group}) + re.NoError(err) + re.Equal("Success!", modifyResp.GetBody()) + modifiedResp, err := leaderClient.GetResourceGroup(ctx, getReq) re.NoError(err) re.NotNil(modifiedResp.GetGroup()) diff --git a/tests/server/tso/tso_proxy_test.go b/tests/server/tso/tso_proxy_test.go index 41508e6efba..95a11115f6d 100644 --- a/tests/server/tso/tso_proxy_test.go +++ b/tests/server/tso/tso_proxy_test.go @@ -146,9 +146,15 @@ func (s *tsoProxyTestSuite) TestRejectFollowerForwardedHost() { client, conn := testutil.MustNewGrpcClient(re, s.leader.GetAddr()) defer conn.Close() + s.verifyForwardedHostRejected(client, s.follower.GetAddr()) + s.verifyForwardedHostRejected(s.pdClient, s.follower.GetAddr()) +} + +func (s *tsoProxyTestSuite) verifyForwardedHostRejected(client pdpb.PDClient, forwardedHost string) { + re := s.Require() ctx, cancel := context.WithCancel(context.Background()) defer cancel() - ctx = grpcutil.BuildForwardContext(ctx, s.follower.GetAddr()) + ctx = grpcutil.BuildForwardContext(ctx, forwardedHost) _, err := client.GetAllStores(ctx, &pdpb.GetAllStoresRequest{Header: s.defaultReq.GetHeader()}) re.Error(err) re.Equal(codes.InvalidArgument, status.Code(err)) From 82502879b375fd2db082c20deb77b0721076cfc4 Mon Sep 17 00:00:00 2001 From: Ryan Leung Date: Fri, 31 Jul 2026 14:57:54 +0800 Subject: [PATCH 7/8] tests: preserve direct follower TSO behavior Signed-off-by: Ryan Leung --- tests/server/tso/tso_test.go | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/server/tso/tso_test.go b/tests/server/tso/tso_test.go index d1eaa1ab955..d5421620161 100644 --- a/tests/server/tso/tso_test.go +++ b/tests/server/tso/tso_test.go @@ -88,7 +88,8 @@ func (s *tsoTestSuite) checkRequestFollower(cluster *tests.TestCluster) { Header: testutil.NewRequestHeader(clusterID), Count: 1, } - ctx = grpcutil.BuildForwardContext(ctx, followerServer.GetAddr()) + // Connect directly without forwarding metadata to verify the original + // follower request behavior. tsoClient, err := grpcPDClient.Tso(ctx) re.NoError(err) defer func() { From 8985d2b0ae236d0ab2a40cc257fb96ba475f438c Mon Sep 17 00:00:00 2001 From: Ryan Leung Date: Fri, 31 Jul 2026 16:25:11 +0800 Subject: [PATCH 8/8] tests: cover forwarded heartbeat validation Signed-off-by: Ryan Leung --- tests/server/tso/tso_proxy_test.go | 20 ++++++++++++++++++++ 1 file changed, 20 insertions(+) diff --git a/tests/server/tso/tso_proxy_test.go b/tests/server/tso/tso_proxy_test.go index 95a11115f6d..f10488c854b 100644 --- a/tests/server/tso/tso_proxy_test.go +++ b/tests/server/tso/tso_proxy_test.go @@ -171,6 +171,26 @@ func (s *tsoProxyTestSuite) verifyForwardedHostRejected(client pdpb.PDClient, fo _, err = tsoClient.Recv() re.Error(err) re.Equal(codes.InvalidArgument, status.Code(err)) + + regionHeartbeatClient, err := client.RegionHeartbeat(ctx) + re.NoError(err) + defer func() { + err := regionHeartbeatClient.CloseSend() + if err != nil && err != io.EOF { + re.NoError(err) + } + }() + re.NoError(regionHeartbeatClient.Send(&pdpb.RegionHeartbeatRequest{Header: s.defaultReq.GetHeader()})) + _, err = regionHeartbeatClient.Recv() + re.Error(err) + re.Equal(codes.InvalidArgument, status.Code(err)) + + reportBucketsClient, err := client.ReportBuckets(ctx) + re.NoError(err) + re.NoError(reportBucketsClient.Send(&pdpb.ReportBucketsRequest{Header: s.defaultReq.GetHeader()})) + _, err = reportBucketsClient.CloseAndRecv() + re.Error(err) + re.Equal(codes.InvalidArgument, status.Code(err)) } func (s *tsoProxyTestSuite) TestRejectUnknownForwardedHost() {