diff --git a/server/channels/api4/user.go b/server/channels/api4/user.go index e223d0cf9f3..175ece3d14d 100644 --- a/server/channels/api4/user.go +++ b/server/channels/api4/user.go @@ -130,7 +130,7 @@ func loginSSOCodeExchange(c *Context, w http.ResponseWriter, r *http.Request) { } // Consume one-time code atomically - token, appErr := c.App.ConsumeTokenOnce(loginCode) + token, appErr := c.App.ConsumeTokenOnce(model.TokenTypeSSOCodeExchange, loginCode) if appErr != nil { c.Err = appErr return diff --git a/server/channels/api4/user_test.go b/server/channels/api4/user_test.go index 3aea66b93ec..09b9adda364 100644 --- a/server/channels/api4/user_test.go +++ b/server/channels/api4/user_test.go @@ -6,6 +6,8 @@ package api4 import ( "bytes" "context" + "crypto/sha256" + "encoding/base64" "encoding/json" "fmt" "image/png" @@ -8490,6 +8492,85 @@ func TestLoginWithDesktopToken(t *testing.T) { }) } +func TestLoginSSOCodeExchange(t *testing.T) { + mainHelper.Parallel(t) + th := Setup(t).InitBasic() + defer th.TearDown() + + t.Run("wrong token type cannot be used for code exchange", func(t *testing.T) { + th.App.UpdateConfig(func(cfg *model.Config) { + cfg.FeatureFlags.MobileSSOCodeExchange = true + }) + + token := model.NewToken(model.TokenTypeOAuth, "extra-data") + require.NoError(t, th.App.Srv().Store().Token().Save(token)) + defer func() { + _ = th.App.Srv().Store().Token().Delete(token.Token) + }() + + props := map[string]string{ + "login_code": token.Token, + "code_verifier": "test_verifier", + "state": "test_state", + } + + resp, err := th.Client.DoAPIPost(context.Background(), "/users/login/sso/code-exchange", model.MapToJSON(props)) + require.Error(t, err) + require.Equal(t, http.StatusNotFound, resp.StatusCode) + }) + + t.Run("successful code exchange with S256 challenge", func(t *testing.T) { + th.App.UpdateConfig(func(cfg *model.Config) { + cfg.FeatureFlags.MobileSSOCodeExchange = true + }) + + samlUser := th.CreateUserWithAuth(model.UserAuthServiceSaml) + + codeVerifier := "test_code_verifier_123456789" + state := "test_state_value" + + sum := sha256.Sum256([]byte(codeVerifier)) + codeChallenge := base64.RawURLEncoding.EncodeToString(sum[:]) + + extra := map[string]string{ + "user_id": samlUser.Id, + "code_challenge": codeChallenge, + "code_challenge_method": "S256", + "state": state, + } + + token := model.NewToken(model.TokenTypeSSOCodeExchange, model.MapToJSON(extra)) + require.NoError(t, th.App.Srv().Store().Token().Save(token)) + + props := map[string]string{ + "login_code": token.Token, + "code_verifier": codeVerifier, + "state": state, + } + + resp, err := th.Client.DoAPIPost(context.Background(), "/users/login/sso/code-exchange", model.MapToJSON(props)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, resp.StatusCode) + + var result map[string]string + require.NoError(t, json.NewDecoder(resp.Body).Decode(&result)) + assert.NotEmpty(t, result["token"]) + assert.NotEmpty(t, result["csrf"]) + + _, err = th.App.Srv().Store().Token().GetByToken(token.Token) + require.Error(t, err) + + authenticatedClient := model.NewAPIv4Client(th.Client.URL) + authenticatedClient.SetToken(result["token"]) + + user, _, err := authenticatedClient.GetMe(context.Background(), "") + require.NoError(t, err) + assert.Equal(t, samlUser.Id, user.Id) + assert.Equal(t, samlUser.Email, user.Email) + assert.Equal(t, samlUser.Username, user.Username) + }) +} + func TestGetUsersByNames(t *testing.T) { mainHelper.Parallel(t) th := Setup(t).InitBasic() diff --git a/server/channels/app/oauth.go b/server/channels/app/oauth.go index 592ec0f96a5..48c8e0bd383 100644 --- a/server/channels/app/oauth.go +++ b/server/channels/app/oauth.go @@ -977,7 +977,7 @@ func (a *App) SwitchEmailToOAuth(rctx request.CTX, w http.ResponseWriter, r *htt stateProps["email"] = email if service == model.UserAuthServiceSaml { - samlToken, samlErr := a.CreateSamlRelayToken(email) + samlToken, samlErr := a.CreateSamlRelayToken(model.TokenTypeSaml, email) if samlErr != nil { return "", samlErr } diff --git a/server/channels/app/saml.go b/server/channels/app/saml.go index 4ff020a1159..483af405404 100644 --- a/server/channels/app/saml.go +++ b/server/channels/app/saml.go @@ -298,8 +298,8 @@ func (a *App) ResetSamlAuthDataToEmail(includeDeleted bool, dryRun bool, userIDs return } -func (a *App) CreateSamlRelayToken(extra string) (*model.Token, *model.AppError) { - token := model.NewToken(model.TokenTypeSaml, extra) +func (a *App) CreateSamlRelayToken(tokenType string, extra string) (*model.Token, *model.AppError) { + token := model.NewToken(tokenType, extra) if err := a.Srv().Store().Token().Save(token); err != nil { var appErr *model.AppError diff --git a/server/channels/app/user.go b/server/channels/app/user.go index c00fc37e1f0..e01705c4a92 100644 --- a/server/channels/app/user.go +++ b/server/channels/app/user.go @@ -1750,8 +1750,8 @@ func (a *App) GetTokenById(token string) (*model.Token, *model.AppError) { return rtoken, nil } -func (a *App) ConsumeTokenOnce(tokenStr string) (*model.Token, *model.AppError) { - token, err := a.Srv().Store().Token().ConsumeOnce(tokenStr) +func (a *App) ConsumeTokenOnce(tokenType, tokenStr string) (*model.Token, *model.AppError) { + token, err := a.Srv().Store().Token().ConsumeOnce(tokenType, tokenStr) if err != nil { var status int switch err.(type) { diff --git a/server/channels/app/user_test.go b/server/channels/app/user_test.go index 178a903331e..b6ea9b39a23 100644 --- a/server/channels/app/user_test.go +++ b/server/channels/app/user_test.go @@ -8,6 +8,7 @@ import ( "database/sql" "encoding/json" "errors" + "net/http" "os" "path/filepath" "strings" @@ -2483,3 +2484,84 @@ func TestRemoteUserDirectChannelCreation(t *testing.T) { assert.Equal(t, model.ChannelTypeDirect, channel.Type) }) } + +func TestConsumeTokenOnce(t *testing.T) { + mainHelper.Parallel(t) + th := Setup(t).InitBasic() + defer th.TearDown() + + t.Run("successfully consume valid token", func(t *testing.T) { + token := model.NewToken(model.TokenTypeOAuth, "extra-data") + require.NoError(t, th.App.Srv().Store().Token().Save(token)) + + consumedToken, appErr := th.App.ConsumeTokenOnce(model.TokenTypeOAuth, token.Token) + require.Nil(t, appErr) + require.NotNil(t, consumedToken) + assert.Equal(t, token.Token, consumedToken.Token) + assert.Equal(t, model.TokenTypeOAuth, consumedToken.Type) + assert.Equal(t, "extra-data", consumedToken.Extra) + + _, err := th.App.Srv().Store().Token().GetByToken(token.Token) + require.Error(t, err) + }) + + t.Run("token not found returns 404", func(t *testing.T) { + nonExistentToken := model.NewRandomString(model.TokenSize) + + consumedToken, appErr := th.App.ConsumeTokenOnce(model.TokenTypeOAuth, nonExistentToken) + require.NotNil(t, appErr) + require.Nil(t, consumedToken) + assert.Equal(t, http.StatusNotFound, appErr.StatusCode) + assert.Equal(t, "ConsumeTokenOnce", appErr.Where) + }) + + t.Run("wrong token type returns not found", func(t *testing.T) { + token := model.NewToken(model.TokenTypeOAuth, "extra-data") + require.NoError(t, th.App.Srv().Store().Token().Save(token)) + defer func() { + _ = th.App.Srv().Store().Token().Delete(token.Token) + }() + + consumedToken, appErr := th.App.ConsumeTokenOnce(model.TokenTypeSaml, token.Token) + require.NotNil(t, appErr) + require.Nil(t, consumedToken) + assert.Equal(t, http.StatusNotFound, appErr.StatusCode) + + _, err := th.App.Srv().Store().Token().GetByToken(token.Token) + require.NoError(t, err) + }) + + t.Run("token can only be consumed once", func(t *testing.T) { + token := model.NewToken(model.TokenTypeSSOCodeExchange, "extra-data") + require.NoError(t, th.App.Srv().Store().Token().Save(token)) + + consumedToken1, appErr := th.App.ConsumeTokenOnce(model.TokenTypeSSOCodeExchange, token.Token) + require.Nil(t, appErr) + require.NotNil(t, consumedToken1) + + consumedToken2, appErr := th.App.ConsumeTokenOnce(model.TokenTypeSSOCodeExchange, token.Token) + require.NotNil(t, appErr) + require.Nil(t, consumedToken2) + assert.Equal(t, http.StatusNotFound, appErr.StatusCode) + }) + + t.Run("empty token string returns not found", func(t *testing.T) { + consumedToken, appErr := th.App.ConsumeTokenOnce(model.TokenTypeOAuth, "") + require.NotNil(t, appErr) + require.Nil(t, consumedToken) + assert.Equal(t, http.StatusNotFound, appErr.StatusCode) + }) + + t.Run("empty token type returns not found", func(t *testing.T) { + token := model.NewToken(model.TokenTypeOAuth, "extra-data") + require.NoError(t, th.App.Srv().Store().Token().Save(token)) + defer func() { + _ = th.App.Srv().Store().Token().Delete(token.Token) + }() + + consumedToken, appErr := th.App.ConsumeTokenOnce("", token.Token) + require.NotNil(t, appErr) + require.Nil(t, consumedToken) + assert.Equal(t, http.StatusNotFound, appErr.StatusCode) + }) +} diff --git a/server/channels/store/retrylayer/retrylayer.go b/server/channels/store/retrylayer/retrylayer.go index cff4ba16442..0dbd66b9dd4 100644 --- a/server/channels/store/retrylayer/retrylayer.go +++ b/server/channels/store/retrylayer/retrylayer.go @@ -14293,11 +14293,11 @@ func (s *RetryLayerTokenStore) Cleanup(expiryTime int64) { } -func (s *RetryLayerTokenStore) ConsumeOnce(tokenStr string) (*model.Token, error) { +func (s *RetryLayerTokenStore) ConsumeOnce(tokenType string, tokenStr string) (*model.Token, error) { tries := 0 for { - result, err := s.TokenStore.ConsumeOnce(tokenStr) + result, err := s.TokenStore.ConsumeOnce(tokenType, tokenStr) if err == nil { return result, nil } diff --git a/server/channels/store/sqlstore/tokens_store.go b/server/channels/store/sqlstore/tokens_store.go index 56b5fb6a823..23b24e73d39 100644 --- a/server/channels/store/sqlstore/tokens_store.go +++ b/server/channels/store/sqlstore/tokens_store.go @@ -78,16 +78,16 @@ func (s SqlTokenStore) GetByToken(tokenString string) (*model.Token, error) { return &token, nil } -func (s SqlTokenStore) ConsumeOnce(tokenStr string) (*model.Token, error) { +func (s SqlTokenStore) ConsumeOnce(tokenType, tokenStr string) (*model.Token, error) { var token model.Token - query := `DELETE FROM Tokens WHERE Token = ? RETURNING *` + query := `DELETE FROM Tokens WHERE Type = ? AND Token = ? RETURNING *` - if err := s.GetMaster().Get(&token, query, tokenStr); err != nil { + if err := s.GetMaster().Get(&token, query, tokenType, tokenStr); err != nil { if err == sql.ErrNoRows { return nil, store.NewErrNotFound("Token", tokenStr) } - return nil, errors.Wrapf(err, "failed to consume token") + return nil, errors.Wrapf(err, "failed to consume token with type %s", tokenType) } return &token, nil diff --git a/server/channels/store/store.go b/server/channels/store/store.go index 783b6972ae2..1a68674e7c9 100644 --- a/server/channels/store/store.go +++ b/server/channels/store/store.go @@ -696,7 +696,7 @@ type TokenStore interface { Save(recovery *model.Token) error Delete(token string) error GetByToken(token string) (*model.Token, error) - ConsumeOnce(tokenStr string) (*model.Token, error) + ConsumeOnce(tokenType, tokenStr string) (*model.Token, error) Cleanup(expiryTime int64) GetAllTokensByType(tokenType string) ([]*model.Token, error) RemoveAllTokensByType(tokenType string) error diff --git a/server/channels/store/storetest/mocks/TokenStore.go b/server/channels/store/storetest/mocks/TokenStore.go index 78ff9f73263..8b0769f4e8c 100644 --- a/server/channels/store/storetest/mocks/TokenStore.go +++ b/server/channels/store/storetest/mocks/TokenStore.go @@ -19,9 +19,9 @@ func (_m *TokenStore) Cleanup(expiryTime int64) { _m.Called(expiryTime) } -// ConsumeOnce provides a mock function with given fields: tokenStr -func (_m *TokenStore) ConsumeOnce(tokenStr string) (*model.Token, error) { - ret := _m.Called(tokenStr) +// ConsumeOnce provides a mock function with given fields: tokenType, tokenStr +func (_m *TokenStore) ConsumeOnce(tokenType string, tokenStr string) (*model.Token, error) { + ret := _m.Called(tokenType, tokenStr) if len(ret) == 0 { panic("no return value specified for ConsumeOnce") @@ -29,19 +29,19 @@ func (_m *TokenStore) ConsumeOnce(tokenStr string) (*model.Token, error) { var r0 *model.Token var r1 error - if rf, ok := ret.Get(0).(func(string) (*model.Token, error)); ok { - return rf(tokenStr) + if rf, ok := ret.Get(0).(func(string, string) (*model.Token, error)); ok { + return rf(tokenType, tokenStr) } - if rf, ok := ret.Get(0).(func(string) *model.Token); ok { - r0 = rf(tokenStr) + if rf, ok := ret.Get(0).(func(string, string) *model.Token); ok { + r0 = rf(tokenType, tokenStr) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*model.Token) } } - if rf, ok := ret.Get(1).(func(string) error); ok { - r1 = rf(tokenStr) + if rf, ok := ret.Get(1).(func(string, string) error); ok { + r1 = rf(tokenType, tokenStr) } else { r1 = ret.Error(1) } diff --git a/server/channels/store/storetest/tokens_store.go b/server/channels/store/storetest/tokens_store.go index 5254e8a752e..915e779f94e 100644 --- a/server/channels/store/storetest/tokens_store.go +++ b/server/channels/store/storetest/tokens_store.go @@ -16,6 +16,7 @@ import ( func TestTokensStore(t *testing.T, rctx request.CTX, ss store.Store) { t.Run("TokensCleanup", func(t *testing.T) { testTokensCleanup(t, rctx, ss) }) + t.Run("ConsumeOnce", func(t *testing.T) { testConsumeOnce(t, rctx, ss) }) } func testTokensCleanup(t *testing.T, rctx request.CTX, ss store.Store) { @@ -41,3 +42,130 @@ func testTokensCleanup(t *testing.T, rctx request.CTX, ss store.Store) { require.NoError(t, err) assert.Len(t, tokens, 0) } + +func testConsumeOnce(t *testing.T, rctx request.CTX, ss store.Store) { + t.Run("successfully consume token once", func(t *testing.T) { + token := &model.Token{ + Token: model.NewRandomString(model.TokenSize), + CreateAt: model.GetMillis(), + Type: model.TokenTypeOAuth, + Extra: "test-extra", + } + err := ss.Token().Save(token) + require.NoError(t, err) + + consumedToken, err := ss.Token().ConsumeOnce(model.TokenTypeOAuth, token.Token) + require.NoError(t, err) + assert.Equal(t, token.Token, consumedToken.Token) + assert.Equal(t, token.Type, consumedToken.Type) + assert.Equal(t, token.Extra, consumedToken.Extra) + + tokens, err := ss.Token().GetAllTokensByType(model.TokenTypeOAuth) + require.NoError(t, err) + assert.Len(t, tokens, 0) + }) + + t.Run("second consumption of same token fails", func(t *testing.T) { + token := &model.Token{ + Token: model.NewRandomString(model.TokenSize), + CreateAt: model.GetMillis(), + Type: model.TokenTypeOAuth, + Extra: "test-extra", + } + err := ss.Token().Save(token) + require.NoError(t, err) + + _, err = ss.Token().ConsumeOnce(model.TokenTypeOAuth, token.Token) + require.NoError(t, err) + + _, err = ss.Token().ConsumeOnce(model.TokenTypeOAuth, token.Token) + require.Error(t, err) + var nfErr *store.ErrNotFound + assert.ErrorAs(t, err, &nfErr) + }) + + t.Run("consume with wrong type fails", func(t *testing.T) { + token := &model.Token{ + Token: model.NewRandomString(model.TokenSize), + CreateAt: model.GetMillis(), + Type: model.TokenTypeOAuth, + Extra: "test-extra", + } + err := ss.Token().Save(token) + require.NoError(t, err) + + _, err = ss.Token().ConsumeOnce(model.TokenTypeSSOCodeExchange, token.Token) + require.Error(t, err) + var nfErr *store.ErrNotFound + assert.ErrorAs(t, err, &nfErr) + + tokens, err := ss.Token().GetAllTokensByType(model.TokenTypeOAuth) + require.NoError(t, err) + assert.Len(t, tokens, 1) + + err = ss.Token().Delete(token.Token) + require.NoError(t, err) + }) + + t.Run("consume non-existent token fails", func(t *testing.T) { + nonExistentToken := model.NewRandomString(model.TokenSize) + _, err := ss.Token().ConsumeOnce(model.TokenTypeOAuth, nonExistentToken) + require.Error(t, err) + var nfErr *store.ErrNotFound + assert.ErrorAs(t, err, &nfErr) + }) + + t.Run("multiple tokens with same type can each be consumed once", func(t *testing.T) { + tokens := make([]*model.Token, 3) + for i := range tokens { + tokens[i] = &model.Token{ + Token: model.NewRandomString(model.TokenSize), + CreateAt: model.GetMillis(), + Type: model.TokenTypeOAuth, + Extra: "test-extra", + } + err := ss.Token().Save(tokens[i]) + require.NoError(t, err) + } + + for _, token := range tokens { + consumedToken, err := ss.Token().ConsumeOnce(model.TokenTypeOAuth, token.Token) + require.NoError(t, err) + assert.Equal(t, token.Token, consumedToken.Token) + } + + allTokens, err := ss.Token().GetAllTokensByType(model.TokenTypeOAuth) + require.NoError(t, err) + assert.Len(t, allTokens, 0) + }) + + t.Run("consuming token of different type leaves others intact", func(t *testing.T) { + oauthToken := &model.Token{ + Token: model.NewRandomString(model.TokenSize), + CreateAt: model.GetMillis(), + Type: model.TokenTypeOAuth, + Extra: "oauth-extra", + } + codeExchangeToken := &model.Token{ + Token: model.NewRandomString(model.TokenSize), + CreateAt: model.GetMillis(), + Type: model.TokenTypeSSOCodeExchange, + Extra: "password-extra", + } + err := ss.Token().Save(oauthToken) + require.NoError(t, err) + err = ss.Token().Save(codeExchangeToken) + require.NoError(t, err) + + consumedToken, err := ss.Token().ConsumeOnce(model.TokenTypeOAuth, oauthToken.Token) + require.NoError(t, err) + assert.Equal(t, oauthToken.Token, consumedToken.Token) + + codeExchangeTokens, err := ss.Token().GetAllTokensByType(model.TokenTypeSSOCodeExchange) + require.NoError(t, err) + assert.Len(t, codeExchangeTokens, 1) + + err = ss.Token().Delete(codeExchangeToken.Token) + require.NoError(t, err) + }) +} diff --git a/server/channels/store/timerlayer/timerlayer.go b/server/channels/store/timerlayer/timerlayer.go index b8fed1003c1..732469508f7 100644 --- a/server/channels/store/timerlayer/timerlayer.go +++ b/server/channels/store/timerlayer/timerlayer.go @@ -11243,10 +11243,10 @@ func (s *TimerLayerTokenStore) Cleanup(expiryTime int64) { } } -func (s *TimerLayerTokenStore) ConsumeOnce(tokenStr string) (*model.Token, error) { +func (s *TimerLayerTokenStore) ConsumeOnce(tokenType string, tokenStr string) (*model.Token, error) { start := time.Now() - result, err := s.TokenStore.ConsumeOnce(tokenStr) + result, err := s.TokenStore.ConsumeOnce(tokenType, tokenStr) elapsed := float64(time.Since(start)) / float64(time.Second) if s.Root.Metrics != nil { diff --git a/server/channels/web/saml.go b/server/channels/web/saml.go index de9ef024023..b5cbba94535 100644 --- a/server/channels/web/saml.go +++ b/server/channels/web/saml.go @@ -106,7 +106,7 @@ func completeSaml(c *Context, w http.ResponseWriter, r *http.Request) { return } - //Validate that the user is with SAML and all that + // Validate that the user is with SAML and all that encodedXML := r.FormValue("SAMLResponse") relayState := r.FormValue("RelayState") @@ -161,7 +161,8 @@ func completeSaml(c *Context, w http.ResponseWriter, r *http.Request) { return } - if err = c.App.CheckUserAllAuthenticationCriteria(c.AppContext, user, ""); err != nil { + err = c.App.CheckUserAllAuthenticationCriteria(c.AppContext, user, "") + if err != nil { handleError(err) return } @@ -250,8 +251,10 @@ func completeSaml(c *Context, w http.ResponseWriter, r *http.Request) { "code_challenge": samlChallenge, "code_challenge_method": samlMethod, }) - code := model.NewToken(model.TokenTypeSaml, extra) - if err := c.App.Srv().Store().Token().Save(code); err != nil { + + var code *model.Token + code, err = c.App.CreateSamlRelayToken(model.TokenTypeSSOCodeExchange, extra) + if err != nil { handleError(model.NewAppError("completeSaml", "app.recover.save.app_error", nil, "", http.StatusInternalServerError).Wrap(err)) return } diff --git a/server/public/model/token.go b/server/public/model/token.go index 47171caaced..731f618e540 100644 --- a/server/public/model/token.go +++ b/server/public/model/token.go @@ -8,10 +8,11 @@ import ( ) const ( - TokenSize = 64 - MaxTokenExipryTime = 1000 * 60 * 60 * 48 // 48 hour - TokenTypeOAuth = "oauth" - TokenTypeSaml = "saml" + TokenSize = 64 + MaxTokenExipryTime = 1000 * 60 * 60 * 48 // 48 hour + TokenTypeOAuth = "oauth" + TokenTypeSaml = "saml" + TokenTypeSSOCodeExchange = "sso-code-exchange" ) type Token struct {