mirror of
https://github.com/nais/wonderwall.git
synced 2026-08-23 21:16:14 +00:00
test: move upstream struct to reverseproxy file
This commit is contained in:
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user