Skip to content
Open
Show file tree
Hide file tree
Changes from 5 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
}

// 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) {
leader := s.GetLeader()
if leader == nil || len(leader.GetClientUrls()) == 0 {
return nil, status.Error(codes.Unavailable, "PD leader is not available")
}
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 the PD leader", forwardedHost)
}

func (s *GrpcServer) closeDelegateClient(forwardedHost string) {
client, ok := s.clientConns.LoadAndDelete(forwardedHost)
if !ok {
Expand Down
25 changes: 14 additions & 11 deletions server/grpc_service.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand Down Expand Up @@ -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() {
Expand Down Expand Up @@ -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.getDelegateClient(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
Expand Down Expand Up @@ -1114,7 +1117,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
}
Expand Down Expand Up @@ -1292,7 +1295,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
}
Expand Down
11 changes: 10 additions & 1 deletion server/resource_group_proxy_service.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 != ""
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Outdated
if forwardedHost == "" {
leader := s.GetLeader()
if leader == nil || len(leader.GetClientUrls()) == 0 {
Expand All @@ -72,7 +73,15 @@ func (s *resourceGroupProxyServer) getPDMetadataWriteDelegateClient(ctx context.
if s.isLocalRequest(forwardedHost) {
return nil, "", nil
}
client, err := s.getDelegateClient(ctx, forwardedHost)
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
}
Expand Down
70 changes: 70 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,73 @@ 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, 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))
}
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
Loading