Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 18 additions & 0 deletions server/forward.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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 {
Expand Down
51 changes: 41 additions & 10 deletions server/grpc_service.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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() {
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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 {
Expand All @@ -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]
})
Expand Down Expand Up @@ -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
Expand All @@ -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]
})
Expand Down
6 changes: 5 additions & 1 deletion server/resource_group_proxy_service.go
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
18 changes: 15 additions & 3 deletions tests/integrations/mcs/resourcemanager/redirector_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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())
Expand Down
96 changes: 96 additions & 0 deletions tests/server/tso/tso_proxy_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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))
}
Comment thread
coderabbitai[bot] marked this conversation as resolved.

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()
Expand Down
3 changes: 2 additions & 1 deletion tests/server/tso/tso_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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() {
Expand Down
Loading