From 375ce229f4923205394d8f27925372b2cbf28130 Mon Sep 17 00:00:00 2001 From: JG Heithcock Date: Mon, 6 Oct 2025 15:20:08 -0700 Subject: [PATCH] MM 65084 server-side (#33861) (#34006) (#34044) * MM 65084 server-side (#33861) (#34006) Automatic Merge * Add ConsumeOnce method to store layers --------- Co-authored-by: Mattermost Build --- api/v4/source/users.yaml | 53 ++++++++++ server/channels/api4/user.go | 99 +++++++++++++++++++ server/channels/app/user.go | 15 +++ .../channels/store/retrylayer/retrylayer.go | 21 ++++ .../channels/store/sqlstore/tokens_store.go | 15 +++ server/channels/store/store.go | 1 + .../store/storetest/mocks/TokenStore.go | 30 ++++++ .../channels/store/timerlayer/timerlayer.go | 16 +++ server/channels/web/saml.go | 56 ++++++++++- server/public/model/feature_flags.go | 4 + server/public/model/token.go | 4 + 11 files changed, 309 insertions(+), 5 deletions(-) diff --git a/api/v4/source/users.yaml b/api/v4/source/users.yaml index 504adc820b..4f92e21bff 100644 --- a/api/v4/source/users.yaml +++ b/api/v4/source/users.yaml @@ -74,6 +74,59 @@ $ref: "#/components/responses/Unauthorized" "403": $ref: "#/components/responses/Forbidden" + /api/v4/users/login/sso/code-exchange: + post: + tags: + - users + summary: Exchange SSO login code for session tokens + description: > + Exchange a short-lived login_code for session tokens using SAML code exchange (mobile SSO flow). + This endpoint is part of the mobile SSO code-exchange flow to prevent tokens + from appearing in deep links. + + ##### Permissions + + No permission required. + operationId: LoginSSOCodeExchange + requestBody: + content: + application/json: + schema: + type: object + required: + - login_code + - code_verifier + - state + properties: + login_code: + description: Short-lived one-time code from SSO callback + type: string + code_verifier: + description: SAML verifier to prove code possession + type: string + state: + description: State parameter to prevent CSRF attacks + type: string + description: SSO code exchange object + required: true + responses: + "200": + description: Code exchange successful + content: + application/json: + schema: + type: object + properties: + token: + description: Session token for authentication + type: string + csrf: + description: CSRF token for request validation + type: string + "400": + $ref: "#/components/responses/BadRequest" + "403": + $ref: "#/components/responses/Forbidden" /api/v4/users/logout: post: tags: diff --git a/server/channels/api4/user.go b/server/channels/api4/user.go index 33e398e616..31cc111c4c 100644 --- a/server/channels/api4/user.go +++ b/server/channels/api4/user.go @@ -4,6 +4,8 @@ package api4 import ( + "crypto/sha256" + "encoding/base64" "encoding/json" "fmt" "io" @@ -63,6 +65,7 @@ func (api *API) InitUser() { api.BaseRoutes.User.Handle("/mfa/generate", api.APISessionRequiredMfa(generateMfaSecret)).Methods(http.MethodPost) api.BaseRoutes.Users.Handle("/login", api.APIHandler(login)).Methods(http.MethodPost) + api.BaseRoutes.Users.Handle("/login/sso/code-exchange", api.APIHandler(loginSSOCodeExchange)).Methods(http.MethodPost) api.BaseRoutes.Users.Handle("/login/desktop_token", api.RateLimitedHandler(api.APIHandler(loginWithDesktopToken), model.RateLimitSettings{PerSec: model.NewPointer(2), MaxBurst: model.NewPointer(1)})).Methods(http.MethodPost) api.BaseRoutes.Users.Handle("/login/switch", api.APIHandler(switchAccountType)).Methods(http.MethodPost) api.BaseRoutes.Users.Handle("/login/cws", api.APIHandlerTrustRequester(loginCWS)).Methods(http.MethodPost) @@ -110,6 +113,102 @@ func (api *API) InitUser() { api.BaseRoutes.Users.Handle("/trigger-notify-admin-posts", api.APISessionRequired(handleTriggerNotifyAdminPosts)).Methods(http.MethodPost) } +// loginSSOCodeExchange exchanges a short-lived login_code for session tokens (mobile SAML code exchange) +func loginSSOCodeExchange(c *Context, w http.ResponseWriter, r *http.Request) { + if !c.App.Config().FeatureFlags.MobileSSOCodeExchange { + c.Err = model.NewAppError("loginSSOCodeExchange", "api.oauth.get_access_token.bad_request.app_error", nil, "feature disabled", http.StatusBadRequest) + return + } + props := model.MapFromJSON(r.Body) + loginCode := props["login_code"] + codeVerifier := props["code_verifier"] + state := props["state"] + + if loginCode == "" || codeVerifier == "" || state == "" { + c.SetInvalidParam("login_code | code_verifier | state") + return + } + + // Consume one-time code atomically + token, appErr := c.App.ConsumeTokenOnce(loginCode) + if appErr != nil { + c.Err = appErr + return + } + + // Check token expiration as fallback to cleanup process + if token.IsExpired() { + c.Err = model.NewAppError("loginSSOCodeExchange", "api.oauth.get_access_token.bad_request.app_error", nil, "token expired", http.StatusBadRequest) + return + } + + // Parse extra JSON + extra := model.MapFromJSON(strings.NewReader(token.Extra)) + userID := extra["user_id"] + codeChallenge := extra["code_challenge"] + method := strings.ToUpper(extra["code_challenge_method"]) + expectedState := extra["state"] + + if userID == "" || codeChallenge == "" || expectedState == "" { + c.Err = model.NewAppError("loginSSOCodeExchange", "api.oauth.get_access_token.bad_request.app_error", nil, "", http.StatusBadRequest) + return + } + + if state != expectedState { + c.Err = model.NewAppError("loginSSOCodeExchange", "api.oauth.get_access_token.bad_request.app_error", nil, "state mismatch", http.StatusBadRequest) + return + } + + // Verify SAML challenge + var computed string + switch strings.ToUpper(method) { + case "S256": + sum := sha256.Sum256([]byte(codeVerifier)) + computed = base64.RawURLEncoding.EncodeToString(sum[:]) + case "": + computed = codeVerifier + case "PLAIN": + // Explicitly reject plain method for security + c.Err = model.NewAppError("loginSSOCodeExchange", "api.oauth.get_access_token.bad_request.app_error", nil, "plain SAML challenge method not supported", + http.StatusBadRequest) + return + default: + // Reject unknown methods + c.Err = model.NewAppError("loginSSOCodeExchange", "api.oauth.get_access_token.bad_request.app_error", nil, "unsupported SAML challenge method", http.StatusBadRequest) + return + } + + if computed != codeChallenge { + c.Err = model.NewAppError("loginSSOCodeExchange", "api.oauth.get_access_token.bad_request.app_error", nil, "SAML challenge mismatch", http.StatusBadRequest) + return + } + + // Create session for this user + user, err := c.App.GetUser(userID) + if err != nil { + c.Err = err + return + } + + isMobile := utils.IsMobileRequest(r) + session, err2 := c.App.DoLogin(c.AppContext, w, r, user, "", isMobile, false, true) + if err2 != nil { + c.Err = err2 + return + } + c.AppContext = c.AppContext.WithSession(session) + c.App.AttachSessionCookies(c.AppContext, w, r) + + // Respond with tokens for mobile client to set + resp := map[string]string{ + "token": session.Token, + "csrf": session.GetCSRF(), + } + if err := json.NewEncoder(w).Encode(resp); err != nil { + c.Logger.Warn("Error while writing response", mlog.Err(err)) + } +} + func createUser(c *Context, w http.ResponseWriter, r *http.Request) { var user model.User if jsonErr := json.NewDecoder(r.Body).Decode(&user); jsonErr != nil { diff --git a/server/channels/app/user.go b/server/channels/app/user.go index 5974b0f589..38502acc57 100644 --- a/server/channels/app/user.go +++ b/server/channels/app/user.go @@ -1750,6 +1750,21 @@ 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) + if err != nil { + var status int + switch err.(type) { + case *store.ErrNotFound: + status = http.StatusNotFound + default: + status = http.StatusInternalServerError + } + return nil, model.NewAppError("ConsumeTokenOnce", "api.user.create_user.signup_link_invalid.app_error", nil, "", status).Wrap(err) + } + return token, nil +} + func (a *App) DeleteToken(token *model.Token) *model.AppError { err := a.Srv().Store().Token().Delete(token.Token) if err != nil { diff --git a/server/channels/store/retrylayer/retrylayer.go b/server/channels/store/retrylayer/retrylayer.go index ef5ebaeac8..616c13771f 100644 --- a/server/channels/store/retrylayer/retrylayer.go +++ b/server/channels/store/retrylayer/retrylayer.go @@ -14138,6 +14138,27 @@ func (s *RetryLayerTokenStore) Cleanup(expiryTime int64) { } +func (s *RetryLayerTokenStore) ConsumeOnce(tokenStr string) (*model.Token, error) { + + tries := 0 + for { + result, err := s.TokenStore.ConsumeOnce(tokenStr) + if err == nil { + return result, nil + } + if !isRepeatableError(err) { + return result, err + } + tries++ + if tries >= 3 { + err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") + return result, err + } + timepkg.Sleep(100 * timepkg.Millisecond) + } + +} + func (s *RetryLayerTokenStore) Delete(token string) error { tries := 0 diff --git a/server/channels/store/sqlstore/tokens_store.go b/server/channels/store/sqlstore/tokens_store.go index 9158e528f7..56b5fb6a82 100644 --- a/server/channels/store/sqlstore/tokens_store.go +++ b/server/channels/store/sqlstore/tokens_store.go @@ -78,6 +78,21 @@ func (s SqlTokenStore) GetByToken(tokenString string) (*model.Token, error) { return &token, nil } +func (s SqlTokenStore) ConsumeOnce(tokenStr string) (*model.Token, error) { + var token model.Token + + query := `DELETE FROM Tokens WHERE Token = ? RETURNING *` + + if err := s.GetMaster().Get(&token, query, tokenStr); err != nil { + if err == sql.ErrNoRows { + return nil, store.NewErrNotFound("Token", tokenStr) + } + return nil, errors.Wrapf(err, "failed to consume token") + } + + return &token, nil +} + func (s SqlTokenStore) Cleanup(expiryTime int64) { if _, err := s.GetMaster().Exec("DELETE FROM Tokens WHERE CreateAt < ?", expiryTime); err != nil { mlog.Error("Unable to cleanup token store.") diff --git a/server/channels/store/store.go b/server/channels/store/store.go index 8fdf20d54a..bd9b879658 100644 --- a/server/channels/store/store.go +++ b/server/channels/store/store.go @@ -693,6 +693,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) 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 1dc802299e..78ff9f7326 100644 --- a/server/channels/store/storetest/mocks/TokenStore.go +++ b/server/channels/store/storetest/mocks/TokenStore.go @@ -19,6 +19,36 @@ 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) + + if len(ret) == 0 { + panic("no return value specified for ConsumeOnce") + } + + 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) *model.Token); ok { + r0 = rf(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) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + // Delete provides a mock function with given fields: token func (_m *TokenStore) Delete(token string) error { ret := _m.Called(token) diff --git a/server/channels/store/timerlayer/timerlayer.go b/server/channels/store/timerlayer/timerlayer.go index 608138722c..1f70419ff3 100644 --- a/server/channels/store/timerlayer/timerlayer.go +++ b/server/channels/store/timerlayer/timerlayer.go @@ -11106,6 +11106,22 @@ func (s *TimerLayerTokenStore) Cleanup(expiryTime int64) { } } +func (s *TimerLayerTokenStore) ConsumeOnce(tokenStr string) (*model.Token, error) { + start := time.Now() + + result, err := s.TokenStore.ConsumeOnce(tokenStr) + + elapsed := float64(time.Since(start)) / float64(time.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("TokenStore.ConsumeOnce", success, elapsed) + } + return result, err +} + func (s *TimerLayerTokenStore) Delete(token string) error { start := time.Now() diff --git a/server/channels/web/saml.go b/server/channels/web/saml.go index 57a0e20559..9305537560 100644 --- a/server/channels/web/saml.go +++ b/server/channels/web/saml.go @@ -37,6 +37,10 @@ func loginWithSaml(c *Context, w http.ResponseWriter, r *http.Request) { action := r.URL.Query().Get("action") isMobile := action == model.OAuthActionMobile redirectURL := html.EscapeString(r.URL.Query().Get("redirect_to")) + // Optional SAML challenge parameters for mobile code-exchange + state := r.URL.Query().Get("state") + codeChallenge := r.URL.Query().Get("code_challenge") + codeChallengeMethod := r.URL.Query().Get("code_challenge_method") relayProps := map[string]string{} relayState := "" @@ -61,6 +65,19 @@ func loginWithSaml(c *Context, w http.ResponseWriter, r *http.Request) { relayProps["redirect_to"] = redirectURL } + // Forward SAML challenge values via RelayState so the complete step can prefer code-exchange + if isMobile { + if state != "" { + relayProps["state"] = state + } + if codeChallenge != "" { + relayProps["code_challenge"] = codeChallenge + } + if codeChallengeMethod != "" { + relayProps["code_challenge_method"] = codeChallengeMethod + } + } + desktopToken := r.URL.Query().Get("desktop_token") if desktopToken != "" { relayProps["desktop_token"] = desktopToken @@ -220,7 +237,33 @@ func completeSaml(c *Context, w http.ResponseWriter, r *http.Request) { return } - // If it's not a desktop login we create a session for this SAML User that will be used in their browser or mobile app + // Decide between legacy token-in-URL vs SAML code-exchange for mobile + samlState := relayProps["state"] + samlChallenge := relayProps["code_challenge"] + samlMethod := relayProps["code_challenge_method"] + + if isMobile && hasRedirectURL && samlChallenge != "" && c.App.Config().FeatureFlags.MobileSSOCodeExchange { + // Issue one-time login_code bound to user and SAML challenge values; do not create a session here + extra := model.MapToJSON(map[string]string{ + "user_id": user.Id, + "state": samlState, + "code_challenge": samlChallenge, + "code_challenge_method": samlMethod, + }) + code := model.NewToken(model.TokenTypeSaml, extra) + if err := c.App.Srv().Store().Token().Save(code); err != nil { + handleError(model.NewAppError("completeSaml", "app.recover.save.app_error", nil, "", http.StatusInternalServerError).Wrap(err)) + return + } + + redirectURL = utils.AppendQueryParamsToURL(redirectURL, map[string]string{ + "login_code": code.Token, + }) + utils.RenderMobileAuthComplete(w, redirectURL) + return + } + + // Legacy: create a session and attach tokens (web/mobile without SAML code exchange) session, err := c.App.DoLogin(c.AppContext, w, r, user, "", isMobile, false, true) if err != nil { handleError(err) @@ -235,10 +278,13 @@ func completeSaml(c *Context, w http.ResponseWriter, r *http.Request) { if hasRedirectURL { if isMobile { // Mobile clients with redirect url support - redirectURL = utils.AppendQueryParamsToURL(redirectURL, map[string]string{ - model.SessionCookieToken: c.AppContext.Session().Token, - model.SessionCookieCsrf: c.AppContext.Session().GetCSRF(), - }) + // Legacy mobile path: return tokens only when SAML code exchange was not requested + if samlChallenge == "" { + redirectURL = utils.AppendQueryParamsToURL(redirectURL, map[string]string{ + model.SessionCookieToken: c.AppContext.Session().Token, + model.SessionCookieCsrf: c.AppContext.Session().GetCSRF(), + }) + } utils.RenderMobileAuthComplete(w, redirectURL) } else { http.Redirect(w, r, redirectURL, http.StatusFound) diff --git a/server/public/model/feature_flags.go b/server/public/model/feature_flags.go index 3c8f9d8346..2df593d1f3 100644 --- a/server/public/model/feature_flags.go +++ b/server/public/model/feature_flags.go @@ -70,6 +70,9 @@ type FeatureFlags struct { AttributeBasedAccessControl bool ContentFlagging bool + + // Enable mobile SSO SAML code-exchange flow (no tokens in deep links) + MobileSSOCodeExchange bool } func (f *FeatureFlags) SetDefaults() { @@ -99,6 +102,7 @@ func (f *FeatureFlags) SetDefaults() { f.CustomProfileAttributes = true f.AttributeBasedAccessControl = true f.ContentFlagging = false + f.MobileSSOCodeExchange = true } // ToMap returns the feature flags as a map[string]string diff --git a/server/public/model/token.go b/server/public/model/token.go index c0c73e367a..47171caace 100644 --- a/server/public/model/token.go +++ b/server/public/model/token.go @@ -41,3 +41,7 @@ func (t *Token) IsValid() *AppError { return nil } + +func (t *Token) IsExpired() bool { + return GetMillis() > (t.CreateAt + MaxTokenExipryTime) +}