From 2ca79b595a3af1572fb4557e4b394c58498f582a Mon Sep 17 00:00:00 2001 From: Trong Huu Nguyen Date: Mon, 28 Apr 2025 11:11:16 +0200 Subject: [PATCH] test: move upstream struct to reverseproxy file --- pkg/handler/handler_test.go | 76 ------------------------------- pkg/handler/reverseproxy_test.go | 77 ++++++++++++++++++++++++++++++++ 2 files changed, 77 insertions(+), 76 deletions(-) diff --git a/pkg/handler/handler_test.go b/pkg/handler/handler_test.go index 589e443..7a726ae 100644 --- a/pkg/handler/handler_test.go +++ b/pkg/handler/handler_test.go @@ -6,13 +6,11 @@ import ( "fmt" "io" "net/http" - "net/http/httptest" "net/url" "strings" "testing" "time" - "github.com/lestrrat-go/jwx/v2/jwt" "github.com/stretchr/testify/assert" "github.com/nais/wonderwall/pkg/config" @@ -637,80 +635,6 @@ func body(t *testing.T, resp *http.Response) string { return string(body) } -type upstream struct { - Server *httptest.Server - URL *url.URL - idp *mock.IdentityProvider - reverseProxyURL *url.URL - requestCallback func(r *http.Request) -} - -func (u *upstream) SetIdentityProvider(idp *mock.IdentityProvider) { - u.idp = idp - u.setReverseProxyUrl(idp.RelyingPartyServer.URL) -} - -func (u *upstream) setReverseProxyUrl(raw string) { - parsed, err := url.Parse(raw) - if err != nil { - panic(err) - } - - u.reverseProxyURL = parsed -} - -func (u *upstream) hasValidToken(r *http.Request) bool { - authHeader := r.Header.Get("Authorization") - token := strings.TrimPrefix(authHeader, "Bearer ") - if len(token) <= 0 { - return false - } - - jwks, err := u.idp.ProviderHandler.Provider.GetPublicJwkSet(r.Context()) - if err != nil { - panic(err) - } - - opts := []jwt.ParseOption{ - jwt.WithValidate(true), - jwt.WithKeySet(*jwks), - jwt.WithIssuer(u.idp.OpenIDConfig.Provider().Issuer()), - jwt.WithAudience(u.idp.OpenIDConfig.Client().ClientID()), - } - - _, err = jwt.ParseString(token, opts...) - return err == nil -} - -func newUpstream(t *testing.T) *upstream { - u := new(upstream) - u.requestCallback = func(r *http.Request) {} - - upstreamHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - u.requestCallback(r) - - // Host should match the original authority from the ingress used to reach Wonderwall - assert.Equal(t, u.reverseProxyURL.Host, r.Host) - assert.NotEqual(t, u.URL.Host, r.Host) - - if u.hasValidToken(r) { - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte("ok")) - } else { - w.WriteHeader(http.StatusUnauthorized) - _, _ = w.Write([]byte("not ok")) - } - }) - server := httptest.NewServer(upstreamHandler) - - upstreamURL, err := url.Parse(server.URL) - assert.NoError(t, err) - - u.Server = server - u.URL = upstreamURL - return u -} - func getCookieFromJar(name string, cookies []*http.Cookie) *http.Cookie { for _, c := range cookies { if c.Name == name { diff --git a/pkg/handler/reverseproxy_test.go b/pkg/handler/reverseproxy_test.go index 5897889..6e0000d 100644 --- a/pkg/handler/reverseproxy_test.go +++ b/pkg/handler/reverseproxy_test.go @@ -2,9 +2,12 @@ package handler_test import ( "net/http" + "net/http/httptest" "net/url" + "strings" "testing" + "github.com/lestrrat-go/jwx/v2/jwt" "github.com/stretchr/testify/assert" "github.com/nais/wonderwall/pkg/mock" @@ -505,3 +508,77 @@ func TestReverseProxy(t *testing.T) { assertUpstreamUnauthorizedResponse(t, resp) }) } + +type upstream struct { + Server *httptest.Server + URL *url.URL + idp *mock.IdentityProvider + reverseProxyURL *url.URL + requestCallback func(r *http.Request) +} + +func newUpstream(t *testing.T) *upstream { + u := new(upstream) + u.requestCallback = func(r *http.Request) {} + + upstreamHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + u.requestCallback(r) + + // Host should match the original authority from the ingress used to reach Wonderwall + assert.Equal(t, u.reverseProxyURL.Host, r.Host) + assert.NotEqual(t, u.URL.Host, r.Host) + + if u.hasValidToken(r) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte("ok")) + } else { + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte("not ok")) + } + }) + server := httptest.NewServer(upstreamHandler) + + upstreamURL, err := url.Parse(server.URL) + assert.NoError(t, err) + + u.Server = server + u.URL = upstreamURL + return u +} + +func (u *upstream) SetIdentityProvider(idp *mock.IdentityProvider) { + u.idp = idp + u.setReverseProxyUrl(idp.RelyingPartyServer.URL) +} + +func (u *upstream) setReverseProxyUrl(raw string) { + parsed, err := url.Parse(raw) + if err != nil { + panic(err) + } + + u.reverseProxyURL = parsed +} + +func (u *upstream) hasValidToken(r *http.Request) bool { + authHeader := r.Header.Get("Authorization") + token := strings.TrimPrefix(authHeader, "Bearer ") + if len(token) <= 0 { + return false + } + + jwks, err := u.idp.ProviderHandler.Provider.GetPublicJwkSet(r.Context()) + if err != nil { + panic(err) + } + + opts := []jwt.ParseOption{ + jwt.WithValidate(true), + jwt.WithKeySet(*jwks), + jwt.WithIssuer(u.idp.OpenIDConfig.Provider().Issuer()), + jwt.WithAudience(u.idp.OpenIDConfig.Client().ClientID()), + } + + _, err = jwt.ParseString(token, opts...) + return err == nil +}