diff --git a/pkg/middleware/auth.go b/pkg/middleware/auth.go index 19acd746bb2..c44d7dd9a7c 100644 --- a/pkg/middleware/auth.go +++ b/pkg/middleware/auth.go @@ -2,6 +2,7 @@ package middleware import ( "net/url" + "regexp" "strconv" "strings" @@ -53,18 +54,19 @@ func notAuthorized(c *models.ReqContext) { redirectTo = setting.AppSubUrl + c.Req.RequestURI } - // remove forceLogin query param if it exists - if parsed, err := url.ParseRequestURI(redirectTo); err == nil { - params := parsed.Query() - params.Del("forceLogin") - parsed.RawQuery = params.Encode() - WriteCookie(c.Resp, "redirect_to", url.QueryEscape(parsed.String()), 0, newCookieOptions) - } else { - c.Logger.Debug("Failed parsing request URI; redirect cookie will not be set", "redirectTo", redirectTo, "error", err) - } + // remove any forceLogin=true params + redirectTo = removeForceLoginParams(redirectTo) + + WriteCookie(c.Resp, "redirect_to", url.QueryEscape(redirectTo), 0, newCookieOptions) c.Redirect(setting.AppSubUrl + "/login") } +var forceLoginParamsRegexp = regexp.MustCompile(`&?forceLogin=true`) + +func removeForceLoginParams(str string) string { + return forceLoginParamsRegexp.ReplaceAllString(str, "") +} + func EnsureEditorOrViewerCanEdit(c *models.ReqContext) { if !c.SignedInUser.HasRole(models.ROLE_EDITOR) && !setting.ViewersCanEdit { accessForbidden(c) diff --git a/pkg/middleware/auth_test.go b/pkg/middleware/auth_test.go index ce0c2587cbb..6f6b243d8a9 100644 --- a/pkg/middleware/auth_test.go +++ b/pkg/middleware/auth_test.go @@ -1,11 +1,13 @@ package middleware import ( + "fmt" "testing" "github.com/grafana/grafana/pkg/bus" "github.com/grafana/grafana/pkg/models" "github.com/grafana/grafana/pkg/setting" + "github.com/stretchr/testify/require" . "github.com/smartystreets/goconvey/convey" ) @@ -104,3 +106,22 @@ func TestMiddlewareAuth(t *testing.T) { }) }) } + +func TestRemoveForceLoginparams(t *testing.T) { + tcs := []struct { + inp string + exp string + }{ + {inp: "/?forceLogin=true", exp: "/?"}, + {inp: "/d/dash/dash-title?ordId=1&forceLogin=true", exp: "/d/dash/dash-title?ordId=1"}, + {inp: "/?kiosk&forceLogin=true", exp: "/?kiosk"}, + {inp: "/d/dash/dash-title?ordId=1&kiosk&forceLogin=true", exp: "/d/dash/dash-title?ordId=1&kiosk"}, + {inp: "/d/dash/dash-title?ordId=1&forceLogin=true&kiosk", exp: "/d/dash/dash-title?ordId=1&kiosk"}, + {inp: "/d/dash/dash-title?forceLogin=true&kiosk", exp: "/d/dash/dash-title?&kiosk"}, + } + for i, tc := range tcs { + t.Run(fmt.Sprintf("testcase %d", i), func(t *testing.T) { + require.Equal(t, tc.exp, removeForceLoginParams(tc.inp)) + }) + } +}