diff --git a/server/forward.go b/server/forward.go index a5d97bfe81e..525a6e3c9aa 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,22 @@ func (s *GrpcServer) getDelegateClient(ctx context.Context, forwardedHost string return conn.(*grpc.ClientConn), nil } +// 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 status.Error(codes.Unavailable, "PD leader is not available") + } + for _, clientURL := range leader.GetClientUrls() { + if clientURL == forwardedHost { + return nil + } + } + return status.Errorf(codes.InvalidArgument, "forwarded host %q is not a client URL of the PD leader", forwardedHost) +} + 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..b68106a39ab 100644 --- a/server/grpc_service.go +++ b/server/grpc_service.go @@ -230,6 +230,11 @@ 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.getDelegateClient(ctx, forwardedHost) if err != nil { @@ -504,9 +509,12 @@ 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 + forwardedHostChecked = forwardedHost == "" ) defer func() { @@ -563,14 +571,21 @@ func (s *GrpcServer) Tso(stream pdpb.PD_TsoServer) error { return errs.ErrNotStarted } - forwardedHost := grpcutil.GetForwardedHost(stream.Context()) - if !s.isLocalRequest(forwardedHost) { - clientConn, err := s.getDelegateClient(s.ctx, forwardedHost) - if err != nil { + 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.getDelegateClient(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 @@ -1080,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 { @@ -1105,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] }) @@ -1261,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 @@ -1283,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] }) diff --git a/server/resource_group_proxy_service.go b/server/resource_group_proxy_service.go index 72f1f249aea..d44be436f4e 100644 --- a/server/resource_group_proxy_service.go +++ b/server/resource_group_proxy_service.go @@ -62,7 +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) - 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") 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 5ad69d25fcc..f10488c854b 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,99 @@ 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() + + 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, forwardedHost) + _, 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)) + + 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() { + 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 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) + 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() 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() {