Skip to content

Commit ca4f38a

Browse files
committed
Context.Scheme should validate values taken from header
Backport PR #2953 (d1d8ad3) to `v4`
1 parent 2e527a7 commit ca4f38a

2 files changed

Lines changed: 165 additions & 36 deletions

File tree

context.go

Lines changed: 16 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -272,22 +272,35 @@ func (c *context) IsWebSocket() bool {
272272
return strings.EqualFold(upgrade, "websocket")
273273
}
274274

275+
func isValidProto(proto string) bool {
276+
if proto == "" {
277+
return false
278+
}
279+
for _, p := range []string{"http", "https", "ws", "wss"} {
280+
if strings.EqualFold(proto, p) {
281+
return true
282+
}
283+
}
284+
return false
285+
}
286+
287+
// Scheme returns the HTTP protocol scheme, `http` or `https`.
275288
func (c *context) Scheme() string {
276289
// Can't use `r.Request.URL.Scheme`
277290
// See: https://groups.google.com/forum/#!topic/golang-nuts/pMUkBlQBDF0
278291
if c.IsTLS() {
279292
return "https"
280293
}
281-
if scheme := c.request.Header.Get(HeaderXForwardedProto); scheme != "" {
294+
if scheme := c.request.Header.Get(HeaderXForwardedProto); isValidProto(scheme) {
282295
return scheme
283296
}
284-
if scheme := c.request.Header.Get(HeaderXForwardedProtocol); scheme != "" {
297+
if scheme := c.request.Header.Get(HeaderXForwardedProtocol); isValidProto(scheme) {
285298
return scheme
286299
}
287300
if ssl := c.request.Header.Get(HeaderXForwardedSsl); ssl == "on" {
288301
return "https"
289302
}
290-
if scheme := c.request.Header.Get(HeaderXUrlScheme); scheme != "" {
303+
if scheme := c.request.Header.Get(HeaderXUrlScheme); isValidProto(scheme) {
291304
return scheme
292305
}
293306
return "http"

context_test.go

Lines changed: 149 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -904,60 +904,176 @@ func TestContext_Request(t *testing.T) {
904904
}
905905

906906
func TestContext_Scheme(t *testing.T) {
907-
tests := []struct {
908-
c Context
909-
s string
907+
var testCases = []struct {
908+
name string
909+
givenIsTLS bool
910+
givenHeaders http.Header
911+
expect string
910912
}{
911913
{
912-
&context{
913-
request: &http.Request{
914-
TLS: &tls.ConnectionState{},
915-
},
914+
name: "defaults to http without TLS or headers",
915+
givenIsTLS: false,
916+
givenHeaders: nil,
917+
expect: "http",
918+
},
919+
{
920+
name: "returns https when TLS is enabled",
921+
givenIsTLS: true,
922+
givenHeaders: nil,
923+
expect: "https",
924+
},
925+
{
926+
name: "TLS takes precedence over forwarded proto",
927+
givenIsTLS: true,
928+
givenHeaders: http.Header{
929+
HeaderXForwardedProto: []string{"http"},
916930
},
917-
"https",
931+
expect: "https",
918932
},
919933
{
920-
&context{
921-
request: &http.Request{
922-
Header: http.Header{HeaderXForwardedProto: []string{"https"}},
923-
},
934+
name: "uses X-Forwarded-Proto http",
935+
givenIsTLS: false,
936+
givenHeaders: http.Header{
937+
HeaderXForwardedProto: []string{"http"},
924938
},
925-
"https",
939+
expect: "http",
926940
},
927941
{
928-
&context{
929-
request: &http.Request{
930-
Header: http.Header{HeaderXForwardedProtocol: []string{"http"}},
931-
},
942+
name: "uses X-Forwarded-Proto https",
943+
givenIsTLS: false,
944+
givenHeaders: http.Header{
945+
HeaderXForwardedProto: []string{"https"},
932946
},
933-
"http",
947+
expect: "https",
934948
},
935949
{
936-
&context{
937-
request: &http.Request{
938-
Header: http.Header{HeaderXForwardedSsl: []string{"on"}},
939-
},
950+
name: "X-Forwarded-Proto is case insensitive",
951+
givenIsTLS: false,
952+
givenHeaders: http.Header{
953+
HeaderXForwardedProto: []string{"HTTPS"},
940954
},
941-
"https",
955+
expect: "HTTPS",
942956
},
943957
{
944-
&context{
945-
request: &http.Request{
946-
Header: http.Header{HeaderXUrlScheme: []string{"https"}},
947-
},
958+
name: "uses X-Forwarded-Proto ws",
959+
givenIsTLS: false,
960+
givenHeaders: http.Header{
961+
HeaderXForwardedProto: []string{"ws"},
948962
},
949-
"https",
963+
expect: "ws",
950964
},
951965
{
952-
&context{
953-
request: &http.Request{},
966+
name: "uses X-Forwarded-Proto wss",
967+
givenIsTLS: false,
968+
givenHeaders: http.Header{
969+
HeaderXForwardedProto: []string{"wss"},
970+
},
971+
expect: "wss",
972+
},
973+
{
974+
name: "ignores invalid X-Forwarded-Proto and uses X-Forwarded-Protocol",
975+
givenIsTLS: false,
976+
givenHeaders: http.Header{
977+
HeaderXForwardedProto: []string{"ftp"},
978+
HeaderXForwardedProtocol: []string{"https"},
954979
},
955-
"http",
980+
expect: "https",
981+
},
982+
{
983+
name: "uses X-Forwarded-Protocol",
984+
givenIsTLS: false,
985+
givenHeaders: http.Header{
986+
HeaderXForwardedProtocol: []string{"https"},
987+
},
988+
expect: "https",
989+
},
990+
{
991+
name: "X-Forwarded-Proto takes precedence over X-Forwarded-Protocol",
992+
givenIsTLS: false,
993+
givenHeaders: http.Header{
994+
HeaderXForwardedProto: []string{"http"},
995+
HeaderXForwardedProtocol: []string{"https"},
996+
},
997+
expect: "http",
998+
},
999+
{
1000+
name: "uses X-Forwarded-Ssl on",
1001+
givenIsTLS: false,
1002+
givenHeaders: http.Header{
1003+
HeaderXForwardedSsl: []string{"on"},
1004+
},
1005+
expect: "https",
1006+
},
1007+
{
1008+
name: "X-Forwarded-Ssl on is case sensitive",
1009+
givenIsTLS: false,
1010+
givenHeaders: http.Header{
1011+
HeaderXForwardedSsl: []string{"ON"},
1012+
},
1013+
expect: "http",
1014+
},
1015+
{
1016+
name: "X-Forwarded-Protocol takes precedence over X-Forwarded-Ssl",
1017+
givenIsTLS: false,
1018+
givenHeaders: http.Header{
1019+
HeaderXForwardedProtocol: []string{"http"},
1020+
HeaderXForwardedSsl: []string{"on"},
1021+
},
1022+
expect: "http",
1023+
},
1024+
{
1025+
name: "uses X-Url-Scheme",
1026+
givenIsTLS: false,
1027+
givenHeaders: http.Header{
1028+
HeaderXUrlScheme: []string{"https"},
1029+
},
1030+
expect: "https",
1031+
},
1032+
{
1033+
name: "X-Forwarded-Ssl takes precedence over X-Url-Scheme",
1034+
givenIsTLS: false,
1035+
givenHeaders: http.Header{
1036+
HeaderXForwardedSsl: []string{"on"},
1037+
HeaderXUrlScheme: []string{"http"},
1038+
},
1039+
expect: "https",
1040+
},
1041+
{
1042+
name: "ignores invalid forwarded headers and falls back to http",
1043+
givenIsTLS: false,
1044+
givenHeaders: http.Header{
1045+
HeaderXForwardedProto: []string{"ftp"},
1046+
HeaderXForwardedProtocol: []string{"smtp"},
1047+
HeaderXForwardedSsl: []string{"off"},
1048+
HeaderXUrlScheme: []string{"file"},
1049+
},
1050+
expect: "http",
1051+
},
1052+
{
1053+
name: "ignores empty forwarded proto and uses X-Url-Scheme",
1054+
givenIsTLS: false,
1055+
givenHeaders: http.Header{
1056+
HeaderXForwardedProto: []string{""},
1057+
HeaderXUrlScheme: []string{"https"},
1058+
},
1059+
expect: "https",
9561060
},
9571061
}
9581062

959-
for _, tt := range tests {
960-
assert.Equal(t, tt.s, tt.c.Scheme())
1063+
for _, tc := range testCases {
1064+
t.Run(tc.name, func(t *testing.T) {
1065+
req := httptest.NewRequest(http.MethodGet, "/", nil)
1066+
if tc.givenHeaders != nil {
1067+
req.Header = tc.givenHeaders
1068+
}
1069+
e := New()
1070+
c := e.NewContext(req, nil)
1071+
if tc.givenIsTLS {
1072+
c.Request().TLS = &tls.ConnectionState{}
1073+
}
1074+
1075+
assert.Equal(t, tc.expect, c.Scheme())
1076+
})
9611077
}
9621078
}
9631079

0 commit comments

Comments
 (0)