Skip to content

Commit fa603ec

Browse files
authored
internal/transport: validate balancer and address metadata (#9203)
Fixes #9199 Validate metadata supplied by balancer `PickResult.Metadata` before merging it into the outgoing context, and validate resolver/address metadata before converting it to HTTP/2 header fields. The tests cover invalid address metadata at the transport boundary and invalid balancer metadata through an end-to-end RPC path. RELEASE NOTES: * transport: Invalid metadata supplied by balancers or resolver addresses is rejected before request headers are created.
1 parent 8518a52 commit fa603ec

3 files changed

Lines changed: 142 additions & 0 deletions

File tree

internal/transport/http2_client.go

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -661,6 +661,9 @@ func (t *http2Client) createHeaderFields(ctx context.Context, callHdr *CallHdr)
661661
if isReservedHeader(k) {
662662
continue
663663
}
664+
if err := imetadata.ValidatePair(k, vv...); err != nil {
665+
return nil, status.Error(codes.Internal, err.Error())
666+
}
664667
for _, v := range vv {
665668
headerFields = append(headerFields, hpack.HeaderField{Name: k, Value: encodeMetadataHeader(k, v)})
666669
}

stream.go

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -625,6 +625,9 @@ func (a *csAttempt) newStream() error {
625625
// maintained in it are local to the attempt. When the attempt has to be
626626
// retried, a new instance of csAttempt will be created.
627627
if a.pickResult.Metadata != nil {
628+
if err := imetadata.Validate(a.pickResult.Metadata); err != nil {
629+
return status.Error(codes.Internal, err.Error())
630+
}
628631
// We currently do not have a function it the metadata package which
629632
// merges given metadata with existing metadata in a context. Existing
630633
// function `AppendToOutgoingContext()` takes a variadic argument of key

test/balancer_test.go

Lines changed: 136 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@ import (
2424
"fmt"
2525
"net"
2626
"reflect"
27+
"strings"
2728
"testing"
2829
"time"
2930

@@ -918,6 +919,141 @@ func (s) TestMetadataInPickResult(t *testing.T) {
918919
}
919920
}
920921

922+
type invalidMetadataCCWrapper struct {
923+
balancer.ClientConn
924+
}
925+
926+
func (t *invalidMetadataCCWrapper) UpdateState(state balancer.State) {
927+
state.Picker = &invalidMetadataPicker{picker: state.Picker}
928+
t.ClientConn.UpdateState(state)
929+
}
930+
931+
type invalidMetadataPicker struct {
932+
picker balancer.Picker
933+
}
934+
935+
func (imp *invalidMetadataPicker) Pick(info balancer.PickInfo) (balancer.PickResult, error) {
936+
res, err := imp.picker.Pick(info)
937+
if err != nil {
938+
return balancer.PickResult{}, err
939+
}
940+
res.Metadata = metadata.MD{"bad key": {"value"}}
941+
return res, nil
942+
}
943+
944+
func (s) TestInvalidMetadataInPickResult(t *testing.T) {
945+
ss := &stubserver.StubServer{EmptyCallF: func(context.Context, *testpb.Empty) (*testpb.Empty, error) {
946+
return &testpb.Empty{}, nil
947+
}}
948+
if err := ss.StartServer(); err != nil {
949+
t.Fatalf("Starting test backend: %v", err)
950+
}
951+
defer ss.Stop()
952+
953+
stub.Register(t.Name(), stub.BalancerFuncs{
954+
Init: func(bd *stub.BalancerData) {
955+
cc := &invalidMetadataCCWrapper{ClientConn: bd.ClientConn}
956+
bd.ChildBalancer = balancer.Get(pickfirst.Name).Build(cc, bd.BuildOptions)
957+
},
958+
Close: func(bd *stub.BalancerData) {
959+
bd.ChildBalancer.Close()
960+
},
961+
UpdateClientConnState: func(bd *stub.BalancerData, ccs balancer.ClientConnState) error {
962+
return bd.ChildBalancer.UpdateClientConnState(ccs)
963+
},
964+
})
965+
966+
r := manual.NewBuilderWithScheme("whatever")
967+
r.InitialState(resolver.State{Addresses: []resolver.Address{{Addr: ss.Address}}})
968+
cc, err := grpc.NewClient(r.Scheme()+":///test.server",
969+
grpc.WithTransportCredentials(insecure.NewCredentials()),
970+
grpc.WithResolvers(r),
971+
grpc.WithDefaultServiceConfig(fmt.Sprintf(`{"loadBalancingConfig": [{"%s":{}}]}`, t.Name())),
972+
)
973+
if err != nil {
974+
t.Fatalf("grpc.NewClient(): %v", err)
975+
}
976+
defer cc.Close()
977+
978+
ctx, cancel := context.WithTimeout(context.Background(), defaultTestTimeout)
979+
defer cancel()
980+
wantErr := `header key "bad key" contains illegal characters not in [0-9a-z-_.]`
981+
if _, err := testgrpc.NewTestServiceClient(cc).EmptyCall(ctx, &testpb.Empty{}); status.Code(err) != codes.Internal || !strings.Contains(err.Error(), wantErr) {
982+
t.Fatalf("EmptyCall() error = %v, want code %v and message containing %q", err, codes.Internal, wantErr)
983+
}
984+
}
985+
986+
func (s) TestAddressMetadataValidation(t *testing.T) {
987+
tests := []struct {
988+
name string
989+
md metadata.MD
990+
wantErr string
991+
}{
992+
{
993+
name: "valid",
994+
md: metadata.Pairs("valid-key", "value"),
995+
},
996+
{
997+
name: "invalid key",
998+
md: metadata.Pairs("bad key", "value"),
999+
wantErr: `header key "bad key" contains illegal characters not in [0-9a-z-_.]`,
1000+
},
1001+
{
1002+
name: "invalid value",
1003+
md: metadata.Pairs("valid-key", "bad\x01value"),
1004+
wantErr: `header key "valid-key" contains value with non-printable ASCII characters`,
1005+
},
1006+
}
1007+
1008+
for _, test := range tests {
1009+
t.Run(test.name, func(t *testing.T) {
1010+
testAddressMetadataValidation(t, test.md, test.wantErr)
1011+
})
1012+
}
1013+
}
1014+
1015+
func testAddressMetadataValidation(t *testing.T, addrMD metadata.MD, wantErr string) {
1016+
ss := &stubserver.StubServer{EmptyCallF: func(ctx context.Context, _ *testpb.Empty) (*testpb.Empty, error) {
1017+
if wantErr != "" {
1018+
return nil, status.Error(codes.Internal, "EmptyCall reached backend with invalid address metadata")
1019+
}
1020+
md, _ := metadata.FromIncomingContext(ctx)
1021+
if got := md.Get("valid-key"); !cmp.Equal(got, []string{"value"}) {
1022+
return nil, status.Errorf(codes.Internal, "metadata.Get(\"valid-key\") = %v, want [value]", got)
1023+
}
1024+
return &testpb.Empty{}, nil
1025+
}}
1026+
if err := ss.StartServer(); err != nil {
1027+
t.Fatalf("Starting test backend: %v", err)
1028+
}
1029+
defer ss.Stop()
1030+
1031+
r := manual.NewBuilderWithScheme("whatever")
1032+
addr := imetadata.Set(resolver.Address{Addr: ss.Address}, addrMD)
1033+
r.InitialState(resolver.State{Addresses: []resolver.Address{addr}})
1034+
cc, err := grpc.NewClient(r.Scheme()+":///test.server",
1035+
grpc.WithTransportCredentials(insecure.NewCredentials()),
1036+
grpc.WithResolvers(r),
1037+
)
1038+
if err != nil {
1039+
t.Fatalf("grpc.NewClient(): %v", err)
1040+
}
1041+
defer cc.Close()
1042+
1043+
ctx, cancel := context.WithTimeout(context.Background(), defaultTestTimeout)
1044+
defer cancel()
1045+
_, err = testgrpc.NewTestServiceClient(cc).EmptyCall(ctx, &testpb.Empty{})
1046+
if wantErr == "" {
1047+
if err != nil {
1048+
t.Fatalf("EmptyCall() error = %v, want nil", err)
1049+
}
1050+
return
1051+
}
1052+
if status.Code(err) != codes.Internal || !strings.Contains(err.Error(), wantErr) {
1053+
t.Fatalf("EmptyCall() error = %v, want code %v and message containing %q", err, codes.Internal, wantErr)
1054+
}
1055+
}
1056+
9211057
// TestSubConnShutdown confirms that the Shutdown method on subconns and
9221058
// RemoveSubConn method on ClientConn properly initiates subconn shutdown.
9231059
func (s) TestSubConnShutdown(t *testing.T) {

0 commit comments

Comments
 (0)