diff --git a/api4/saml_test.go b/api4/saml_test.go index 655a0eb9e7..ceae3aab45 100644 --- a/api4/saml_test.go +++ b/api4/saml_test.go @@ -4,7 +4,12 @@ package api4 import ( + "net/http" "testing" + + "github.com/mattermost/mattermost-server/v5/model" + + "github.com/stretchr/testify/require" ) func TestGetSamlMetadata(t *testing.T) { @@ -17,3 +22,31 @@ func TestGetSamlMetadata(t *testing.T) { // Rest is tested by enterprise tests } + +func TestSamlCompleteCSRFPass(t *testing.T) { + th := Setup().InitBasic() + defer th.TearDown() + + url := th.Client.Url + "/login/sso/saml" + req, err := http.NewRequest("POST", url, nil) + if err != nil { + return + } + + cookie1 := &http.Cookie{ + Name: model.SESSION_COOKIE_USER, + Value: th.BasicUser.Username, + } + cookie2 := &http.Cookie{ + Name: model.SESSION_COOKIE_TOKEN, + Value: th.Client.AuthToken, + } + req.AddCookie(cookie1) + req.AddCookie(cookie2) + + client := &http.Client{} + resp, err := client.Do(req) + require.NoError(t, err) + require.NotEqual(t, http.StatusUnauthorized, resp.StatusCode) + defer resp.Body.Close() +} diff --git a/web/saml.go b/web/saml.go index 991dd1807a..e7bd23b2f7 100644 --- a/web/saml.go +++ b/web/saml.go @@ -13,8 +13,8 @@ import ( ) func (w *Web) InitSaml() { - w.MainRouter.Handle("/login/sso/saml", w.NewHandler(loginWithSaml)).Methods("GET") - w.MainRouter.Handle("/login/sso/saml", w.NewHandler(completeSaml)).Methods("POST") + w.MainRouter.Handle("/login/sso/saml", w.ApiHandler(loginWithSaml)).Methods("GET") + w.MainRouter.Handle("/login/sso/saml", w.ApiHandlerTrustRequester(completeSaml)).Methods("POST") } func loginWithSaml(c *Context, w http.ResponseWriter, r *http.Request) {