diff --git a/internal/utilities/request.go b/internal/utilities/request.go index f9988745ea..9115e0a475 100644 --- a/internal/utilities/request.go +++ b/internal/utilities/request.go @@ -140,13 +140,19 @@ func IsRedirectURLValid(config *conf.GlobalConfiguration, redirectURL string) bo // getRedirectTo tries extract redirect url from header or from query params func getRedirectTo(r *http.Request) (reqref string) { - reqref = r.Header.Get("redirect_to") + reqref = r.Header.Get("redirect-to") + if reqref == "" { + reqref = r.Header.Get("redirect_to") + } if reqref != "" { return } if err := r.ParseForm(); err == nil { - reqref = r.Form.Get("redirect_to") + reqref = r.Form.Get("redirect-to") + if reqref == "" { + reqref = r.Form.Get("redirect_to") + } } return diff --git a/internal/utilities/request_test.go b/internal/utilities/request_test.go index b7b825fbbf..ff8e92b0da 100644 --- a/internal/utilities/request_test.go +++ b/internal/utilities/request_test.go @@ -333,6 +333,16 @@ func TestGetReferrer(t *tst.T) { r := httptest.NewRequest("GET", "http://localhost?redirect_to="+c.redirectURL, nil) referrer := GetReferrer(r, &config) require.Equal(t, c.expected, referrer) + + r2 := httptest.NewRequest("GET", "http://localhost", nil) + r2.Header.Set("redirect-to", c.redirectURL) + referrer2 := GetReferrer(r2, &config) + require.Equal(t, c.expected, referrer2) + + r3 := httptest.NewRequest("GET", "http://localhost", nil) + r3.Header.Set("redirect_to", c.redirectURL) + referrer3 := GetReferrer(r3, &config) + require.Equal(t, c.expected, referrer3) }) } }