package reqtls_test import ( "context" "crypto/tls" "net/http" "net/http/httptest" "testing" "github.com/stretchr/testify/assert" "sneak.berlin/go/webhooker/internal/reqtls" ) // newReq builds a plaintext request with no forwarding headers. func newReq(t *testing.T) *http.Request { t.Helper() return httptest.NewRequestWithContext( context.Background(), http.MethodGet, "/", nil, ) } func TestIsTLS_DirectTLS(t *testing.T) { t.Parallel() r := newReq(t) r.TLS = &tls.ConnectionState{} assert.True( t, reqtls.IsTLS(r), "a request that arrived over TLS is TLS", ) } func TestIsTLS_PlaintextNoHeader(t *testing.T) { t.Parallel() assert.False( t, reqtls.IsTLS(newReq(t)), "no TLS connection and no header means plaintext", ) } // protoCase is one X-Forwarded-Proto spelling and the answer IsTLS // owes it. type protoCase struct { name string header string want bool why string } // protoCases enumerates the header values real infrastructure emits. func protoCases() []protoCase { return append(protoTLSCases(), protoPlaintextCases()...) } // protoTLSCases are the spellings that name a TLS client connection. // Every one but the first is a spelling an exact == "https" // comparison used to miss, silently downgrading a genuinely-HTTPS // deployment to the plaintext path. func protoTLSCases() []protoCase { return []protoCase{ { name: "lowercase", header: "https", want: true, why: "the ordinary spelling", }, { name: "uppercase", header: "HTTPS", want: true, why: "the value is a case-insensitive token; " + "nothing obliges a proxy to lowercase it", }, { name: "mixed case", header: "HttpS", want: true, why: "case folding must be total, not just the two extremes", }, { name: "chain with plaintext inner hop", header: "https, http", want: true, why: "a chained proxy appends its hop; the leftmost " + "element is the client-facing one", }, { name: "chain of two TLS hops", header: "https,https", want: true, why: "appended chain with no space after the comma", }, { name: "trailing space", header: "https ", want: true, why: "surrounding whitespace is not part of the token", }, { name: "leading space", header: " https", want: true, why: "surrounding whitespace is not part of the token", }, { name: "uppercase chain", header: "HTTPS, HTTP", want: true, why: "case folding and chain splitting must compose", }, } } // protoPlaintextCases are the values that must NOT be read as TLS. func protoPlaintextCases() []protoCase { return []protoCase{ { name: "plaintext", header: "http", want: false, why: "the negative control: the proxy reports plaintext", }, { name: "plaintext chain with TLS inner hop", header: "http, https", want: false, why: "the client-facing hop is plaintext even though " + "an inner hop used TLS", }, { name: "empty", header: "", want: false, why: "an empty header asserts nothing", }, { name: "whitespace only", header: " ", want: false, why: "a blank header asserts nothing", }, { name: "unrelated token", header: "ftp", want: false, why: "only https means TLS", }, { name: "https as a substring", header: "nothttps", want: false, why: "matching must be on the whole token, not a substring", }, } } func TestIsTLS_ForwardedProtoSpellings(t *testing.T) { t.Parallel() for _, tc := range protoCases() { t.Run(tc.name, func(t *testing.T) { t.Parallel() r := newReq(t) r.Header.Set("X-Forwarded-Proto", tc.header) assert.Equal( t, tc.want, reqtls.IsTLS(r), "X-Forwarded-Proto %q: %s", tc.header, tc.why, ) }) } } // TestIsTLS_DirectTLSBeatsPlaintextHeader pins the precedence: a // connection this process itself terminated with TLS is a fact, and a // header claiming otherwise does not override it. func TestIsTLS_DirectTLSBeatsPlaintextHeader(t *testing.T) { t.Parallel() r := newReq(t) r.TLS = &tls.ConnectionState{} r.Header.Set("X-Forwarded-Proto", "http") assert.True( t, reqtls.IsTLS(r), "an actual TLS connection outranks a header claiming plaintext", ) } // TestIsTLS_FirstHeaderValueWins covers a proxy that adds a second // header line rather than appending to the existing one. net/http // keeps them as separate values; the first is the client-facing hop, // matching how the comma-separated form is read. func TestIsTLS_FirstHeaderValueWins(t *testing.T) { t.Parallel() r := newReq(t) r.Header.Add("X-Forwarded-Proto", "https") r.Header.Add("X-Forwarded-Proto", "http") assert.True( t, reqtls.IsTLS(r), "the first header line is the client-facing hop", ) }