[MM-64896][MM-64898] Pass inviteid/tokenid to relay state/props for external auth when auto-joining a team (#33545) (#33666)

Automatic Merge
Этот коммит содержится в:
Mattermost Build
2025-08-13 13:04:00 +03:00
коммит произвёл GitHub
родитель ed9e2dbbce
Коммит 2d5cdc6e21
7 изменённых файлов: 129 добавлений и 106 удалений

Просмотреть файл

@@ -438,11 +438,14 @@ func (a *App) newSessionUpdateToken(c request.CTX, app *model.OAuthApp, accessDa
return accessRsp, nil return accessRsp, nil
} }
func (a *App) GetOAuthLoginEndpoint(c request.CTX, w http.ResponseWriter, r *http.Request, service, teamID, action, redirectTo, loginHint string, isMobile bool, desktopToken string) (string, *model.AppError) { func (a *App) GetOAuthLoginEndpoint(c request.CTX, w http.ResponseWriter, r *http.Request, service, action, redirectTo, loginHint string, isMobile bool, desktopToken string, inviteToken string, inviteId string) (string, *model.AppError) {
stateProps := map[string]string{} stateProps := map[string]string{}
stateProps["action"] = action stateProps["action"] = action
if teamID != "" {
stateProps["team_id"] = teamID if inviteToken != "" {
stateProps["invite_token"] = inviteToken
} else if inviteId != "" {
stateProps["invite_id"] = inviteId
} }
if redirectTo != "" { if redirectTo != "" {
@@ -463,11 +466,14 @@ func (a *App) GetOAuthLoginEndpoint(c request.CTX, w http.ResponseWriter, r *htt
return authURL, nil return authURL, nil
} }
func (a *App) GetOAuthSignupEndpoint(c request.CTX, w http.ResponseWriter, r *http.Request, service, teamID string, desktopToken string) (string, *model.AppError) { func (a *App) GetOAuthSignupEndpoint(c request.CTX, w http.ResponseWriter, r *http.Request, service, desktopToken string, inviteToken string, inviteId string) (string, *model.AppError) {
stateProps := map[string]string{} stateProps := map[string]string{}
stateProps["action"] = model.OAuthActionSignup stateProps["action"] = model.OAuthActionSignup
if teamID != "" {
stateProps["team_id"] = teamID if inviteToken != "" {
stateProps["invite_token"] = inviteToken
} else if inviteId != "" {
stateProps["invite_id"] = inviteId
} }
if desktopToken != "" { if desktopToken != "" {
@@ -570,22 +576,26 @@ func (a *App) RevokeAccessToken(c request.CTX, token string) *model.AppError {
return nil return nil
} }
func (a *App) CompleteOAuth(c request.CTX, service string, body io.ReadCloser, teamID string, props map[string]string, tokenUser *model.User) (*model.User, *model.AppError) { func (a *App) CompleteOAuth(c request.CTX, service string, body io.ReadCloser, props map[string]string, tokenUser *model.User) (*model.User, *model.AppError) {
defer body.Close() defer body.Close()
action := props["action"] action := props["action"]
// Extract invite token or ID from props so we can add the user to the team if needed
inviteToken := props["invite_token"]
inviteId := props["invite_id"]
switch action { switch action {
case model.OAuthActionSignup: case model.OAuthActionSignup:
return a.CreateOAuthUser(c, service, body, teamID, tokenUser) return a.CreateOAuthUser(c, service, body, inviteToken, inviteId, tokenUser)
case model.OAuthActionLogin: case model.OAuthActionLogin:
return a.LoginByOAuth(c, service, body, teamID, tokenUser) return a.LoginByOAuth(c, service, body, inviteToken, inviteId, tokenUser)
case model.OAuthActionEmailToSSO: case model.OAuthActionEmailToSSO:
return a.CompleteSwitchWithOAuth(c, service, body, props["email"], tokenUser) return a.CompleteSwitchWithOAuth(c, service, body, props["email"], tokenUser)
case model.OAuthActionSSOToEmail: case model.OAuthActionSSOToEmail:
return a.LoginByOAuth(c, service, body, teamID, tokenUser) return a.LoginByOAuth(c, service, body, inviteToken, inviteId, tokenUser)
default: default:
return a.LoginByOAuth(c, service, body, teamID, tokenUser) return a.LoginByOAuth(c, service, body, inviteToken, inviteId, tokenUser)
} }
} }
@@ -606,7 +616,7 @@ func (a *App) getSSOProvider(service string) (einterfaces.OAuthProvider, *model.
return provider, nil return provider, nil
} }
func (a *App) LoginByOAuth(c request.CTX, service string, userData io.Reader, teamID string, tokenUser *model.User) (*model.User, *model.AppError) { func (a *App) LoginByOAuth(c request.CTX, service string, userData io.Reader, inviteToken string, inviteId string, tokenUser *model.User) (*model.User, *model.AppError) {
provider, e := a.getSSOProvider(service) provider, e := a.getSSOProvider(service)
if e != nil { if e != nil {
return nil, e return nil, e
@@ -632,7 +642,7 @@ func (a *App) LoginByOAuth(c request.CTX, service string, userData io.Reader, te
user, err := a.GetUserByAuth(model.NewPointer(*authUser.AuthData), service) user, err := a.GetUserByAuth(model.NewPointer(*authUser.AuthData), service)
if err != nil { if err != nil {
if err.Id == MissingAuthAccountError { if err.Id == MissingAuthAccountError {
user, err = a.CreateOAuthUser(c, service, bytes.NewReader(buf.Bytes()), teamID, tokenUser) user, err = a.CreateOAuthUser(c, service, bytes.NewReader(buf.Bytes()), inviteToken, inviteId, tokenUser)
} else { } else {
return nil, err return nil, err
} }
@@ -647,8 +657,9 @@ func (a *App) LoginByOAuth(c request.CTX, service string, userData io.Reader, te
if err = a.UpdateOAuthUserAttrs(c, bytes.NewReader(buf.Bytes()), user, provider, service, tokenUser); err != nil { if err = a.UpdateOAuthUserAttrs(c, bytes.NewReader(buf.Bytes()), user, provider, service, tokenUser); err != nil {
return nil, err return nil, err
} }
if teamID != "" {
err = a.AddUserToTeamByTeamId(c, teamID, user) if err = a.AddUserToTeamByInviteIfNeeded(c, user, inviteToken, inviteId); err != nil {
c.Logger().Warn("Failed to add user to team", mlog.Err(err))
} }
} }
@@ -802,20 +813,20 @@ func (a *App) GetAuthorizationCode(c request.CTX, w http.ResponseWriter, r *http
return authURL, nil return authURL, nil
} }
func (a *App) AuthorizeOAuthUser(c request.CTX, w http.ResponseWriter, r *http.Request, service, code, state, redirectURI string) (io.ReadCloser, string, map[string]string, *model.User, *model.AppError) { func (a *App) AuthorizeOAuthUser(c request.CTX, w http.ResponseWriter, r *http.Request, service, code, state, redirectURI string) (io.ReadCloser, map[string]string, *model.User, *model.AppError) {
provider, e := a.getSSOProvider(service) provider, e := a.getSSOProvider(service)
if e != nil { if e != nil {
return nil, "", nil, nil, e return nil, nil, nil, e
} }
sso, e2 := provider.GetSSOSettings(c, a.Config(), service) sso, e2 := provider.GetSSOSettings(c, a.Config(), service)
if e2 != nil { if e2 != nil {
return nil, "", nil, nil, model.NewAppError("AuthorizeOAuthUser.GetSSOSettings", "api.user.get_authorization_code.endpoint.app_error", nil, "", http.StatusNotImplemented).Wrap(e2) return nil, nil, nil, model.NewAppError("AuthorizeOAuthUser.GetSSOSettings", "api.user.get_authorization_code.endpoint.app_error", nil, "", http.StatusNotImplemented).Wrap(e2)
} }
b, strErr := b64.StdEncoding.DecodeString(state) b, strErr := b64.StdEncoding.DecodeString(state)
if strErr != nil { if strErr != nil {
return nil, "", nil, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.invalid_state.app_error", nil, "", http.StatusBadRequest).Wrap(strErr) return nil, nil, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.invalid_state.app_error", nil, "", http.StatusBadRequest).Wrap(strErr)
} }
stateStr := string(b) stateStr := string(b)
@@ -823,25 +834,25 @@ func (a *App) AuthorizeOAuthUser(c request.CTX, w http.ResponseWriter, r *http.R
expectedToken, appErr := a.GetOAuthStateToken(stateProps["token"]) expectedToken, appErr := a.GetOAuthStateToken(stateProps["token"])
if appErr != nil { if appErr != nil {
return nil, "", stateProps, nil, appErr return nil, stateProps, nil, appErr
} }
stateEmail := stateProps["email"] stateEmail := stateProps["email"]
stateAction := stateProps["action"] stateAction := stateProps["action"]
if stateAction == model.OAuthActionEmailToSSO && stateEmail == "" { if stateAction == model.OAuthActionEmailToSSO && stateEmail == "" {
err := errors.New("No email provided in state when trying to switch from email to SSO") err := errors.New("No email provided in state when trying to switch from email to SSO")
return nil, "", stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.invalid_state.app_error", nil, "", http.StatusBadRequest).Wrap(err) return nil, stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.invalid_state.app_error", nil, "", http.StatusBadRequest).Wrap(err)
} }
cookie, cookieErr := r.Cookie(CookieOAuth) cookie, cookieErr := r.Cookie(CookieOAuth)
if cookieErr != nil { if cookieErr != nil {
return nil, "", stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.invalid_state.app_error", nil, "", http.StatusBadRequest).Wrap(cookieErr) return nil, stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.invalid_state.app_error", nil, "", http.StatusBadRequest).Wrap(cookieErr)
} }
expectedTokenExtra := generateOAuthStateTokenExtra(stateEmail, stateAction, cookie.Value) expectedTokenExtra := generateOAuthStateTokenExtra(stateEmail, stateAction, cookie.Value)
if expectedTokenExtra != expectedToken.Extra { if expectedTokenExtra != expectedToken.Extra {
err := errors.New("Extra token value does not match token generated from state") err := errors.New("Extra token value does not match token generated from state")
return nil, "", stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.invalid_state.app_error", nil, "", http.StatusBadRequest).Wrap(err) return nil, stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.invalid_state.app_error", nil, "", http.StatusBadRequest).Wrap(err)
} }
appErr = a.DeleteToken(expectedToken) appErr = a.DeleteToken(expectedToken)
@@ -861,8 +872,6 @@ func (a *App) AuthorizeOAuthUser(c request.CTX, w http.ResponseWriter, r *http.R
http.SetCookie(w, httpCookie) http.SetCookie(w, httpCookie)
teamID := stateProps["team_id"]
p := url.Values{} p := url.Values{}
p.Set("client_id", *sso.Id) p.Set("client_id", *sso.Id)
p.Set("client_secret", *sso.Secret) p.Set("client_secret", *sso.Secret)
@@ -872,7 +881,7 @@ func (a *App) AuthorizeOAuthUser(c request.CTX, w http.ResponseWriter, r *http.R
req, requestErr := http.NewRequest("POST", *sso.TokenEndpoint, strings.NewReader(p.Encode())) req, requestErr := http.NewRequest("POST", *sso.TokenEndpoint, strings.NewReader(p.Encode()))
if requestErr != nil { if requestErr != nil {
return nil, "", stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.token_failed.app_error", nil, "", http.StatusInternalServerError).Wrap(requestErr) return nil, stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.token_failed.app_error", nil, "", http.StatusInternalServerError).Wrap(requestErr)
} }
req.Header.Set("Content-Type", "application/x-www-form-urlencoded") req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
@@ -880,7 +889,7 @@ func (a *App) AuthorizeOAuthUser(c request.CTX, w http.ResponseWriter, r *http.R
resp, err := a.HTTPService().MakeClient(true).Do(req) resp, err := a.HTTPService().MakeClient(true).Do(req)
if err != nil { if err != nil {
return nil, "", stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.token_failed.app_error", nil, "", http.StatusInternalServerError).Wrap(err) return nil, stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.token_failed.app_error", nil, "", http.StatusInternalServerError).Wrap(err)
} }
defer resp.Body.Close() defer resp.Body.Close()
@@ -889,15 +898,15 @@ func (a *App) AuthorizeOAuthUser(c request.CTX, w http.ResponseWriter, r *http.R
var ar *model.AccessResponse var ar *model.AccessResponse
err = json.NewDecoder(tee).Decode(&ar) err = json.NewDecoder(tee).Decode(&ar)
if err != nil || resp.StatusCode != http.StatusOK { if err != nil || resp.StatusCode != http.StatusOK {
return nil, "", stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.bad_response.app_error", nil, fmt.Sprintf("response_body=%s, status_code=%d, error=%v", buf.String(), resp.StatusCode, err), http.StatusInternalServerError).Wrap(err) return nil, stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.bad_response.app_error", nil, fmt.Sprintf("response_body=%s, status_code=%d, error=%v", buf.String(), resp.StatusCode, err), http.StatusInternalServerError).Wrap(err)
} }
if strings.ToLower(ar.TokenType) != model.AccessTokenType { if strings.ToLower(ar.TokenType) != model.AccessTokenType {
return nil, "", stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.bad_token.app_error", nil, "token_type="+ar.TokenType+", response_body="+buf.String(), http.StatusInternalServerError) return nil, stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.bad_token.app_error", nil, "token_type="+ar.TokenType+", response_body="+buf.String(), http.StatusInternalServerError)
} }
if ar.AccessToken == "" { if ar.AccessToken == "" {
return nil, "", stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.missing.app_error", nil, "response_body="+buf.String(), http.StatusInternalServerError) return nil, stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.missing.app_error", nil, "response_body="+buf.String(), http.StatusInternalServerError)
} }
p = url.Values{} p = url.Values{}
@@ -907,13 +916,13 @@ func (a *App) AuthorizeOAuthUser(c request.CTX, w http.ResponseWriter, r *http.R
if ar.IdToken != "" { if ar.IdToken != "" {
userFromToken, err = provider.GetUserFromIdToken(c, ar.IdToken) userFromToken, err = provider.GetUserFromIdToken(c, ar.IdToken)
if err != nil { if err != nil {
return nil, "", stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.token_failed.app_error", nil, "", http.StatusInternalServerError).Wrap(err) return nil, stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.token_failed.app_error", nil, "", http.StatusInternalServerError).Wrap(err)
} }
} }
req, requestErr = http.NewRequest("GET", *sso.UserAPIEndpoint, strings.NewReader("")) req, requestErr = http.NewRequest("GET", *sso.UserAPIEndpoint, strings.NewReader(""))
if requestErr != nil { if requestErr != nil {
return nil, "", stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.service.app_error", map[string]any{"Service": service}, "", http.StatusInternalServerError).Wrap(requestErr) return nil, stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.service.app_error", map[string]any{"Service": service}, "", http.StatusInternalServerError).Wrap(requestErr)
} }
req.Header.Set("Content-Type", "application/x-www-form-urlencoded") req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
@@ -922,7 +931,7 @@ func (a *App) AuthorizeOAuthUser(c request.CTX, w http.ResponseWriter, r *http.R
resp, err = a.HTTPService().MakeClient(true).Do(req) resp, err = a.HTTPService().MakeClient(true).Do(req)
if err != nil { if err != nil {
return nil, "", stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.service.app_error", map[string]any{"Service": service}, "", http.StatusInternalServerError).Wrap(err) return nil, stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.service.app_error", map[string]any{"Service": service}, "", http.StatusInternalServerError).Wrap(err)
} else if resp.StatusCode != http.StatusOK { } else if resp.StatusCode != http.StatusOK {
defer resp.Body.Close() defer resp.Body.Close()
@@ -935,17 +944,17 @@ func (a *App) AuthorizeOAuthUser(c request.CTX, w http.ResponseWriter, r *http.R
if service == model.ServiceGitlab && resp.StatusCode == http.StatusForbidden && strings.Contains(bodyString, "Terms of Service") { if service == model.ServiceGitlab && resp.StatusCode == http.StatusForbidden && strings.Contains(bodyString, "Terms of Service") {
url, err := url.Parse(*sso.UserAPIEndpoint) url, err := url.Parse(*sso.UserAPIEndpoint)
if err != nil { if err != nil {
return nil, "", stateProps, nil, model.NewAppError("AuthorizeOAuthUser", model.NoTranslation, nil, "", http.StatusInternalServerError).Wrap(errors.Wrapf(err, "error parsing %s", *sso.UserAPIEndpoint)) return nil, stateProps, nil, model.NewAppError("AuthorizeOAuthUser", model.NoTranslation, nil, "", http.StatusInternalServerError).Wrap(errors.Wrapf(err, "error parsing %s", *sso.UserAPIEndpoint))
} }
// Return a nicer error when the user hasn't accepted GitLab's terms of service // Return a nicer error when the user hasn't accepted GitLab's terms of service
return nil, "", stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "oauth.gitlab.tos.error", map[string]any{"URL": url.Hostname()}, "", http.StatusBadRequest) return nil, stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "oauth.gitlab.tos.error", map[string]any{"URL": url.Hostname()}, "", http.StatusBadRequest)
} }
return nil, "", stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.response.app_error", nil, "response_body="+bodyString, http.StatusInternalServerError) return nil, stateProps, nil, model.NewAppError("AuthorizeOAuthUser", "api.user.authorize_oauth_user.response.app_error", nil, "response_body="+bodyString, http.StatusInternalServerError)
} }
// Note that resp.Body is not closed here, so it must be closed by the caller // Note that resp.Body is not closed here, so it must be closed by the caller
return resp.Body, teamID, stateProps, userFromToken, nil return resp.Body, stateProps, userFromToken, nil
} }
func (a *App) SwitchEmailToOAuth(c request.CTX, w http.ResponseWriter, r *http.Request, email, password, code, service string) (string, *model.AppError) { func (a *App) SwitchEmailToOAuth(c request.CTX, w http.ResponseWriter, r *http.Request, email, password, code, service string) (string, *model.AppError) {

Просмотреть файл

@@ -193,7 +193,7 @@ func TestAuthorizeOAuthUser(t *testing.T) {
th := setup(t, false, true, true, "") th := setup(t, false, true, true, "")
defer th.TearDown() defer th.TearDown()
_, _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, nil, nil, model.ServiceGitlab, "", "", "") _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, nil, nil, model.ServiceGitlab, "", "", "")
require.NotNil(t, err) require.NotNil(t, err)
assert.Equal(t, "api.user.authorize_oauth_user.unsupported.app_error", err.Id) assert.Equal(t, "api.user.authorize_oauth_user.unsupported.app_error", err.Id)
}) })
@@ -204,7 +204,7 @@ func TestAuthorizeOAuthUser(t *testing.T) {
state := "!" state := "!"
_, _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, nil, nil, model.ServiceGitlab, "", state, "") _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, nil, nil, model.ServiceGitlab, "", state, "")
require.NotNil(t, err) require.NotNil(t, err)
assert.Equal(t, "api.user.authorize_oauth_user.invalid_state.app_error", err.Id) assert.Equal(t, "api.user.authorize_oauth_user.invalid_state.app_error", err.Id)
}) })
@@ -217,7 +217,7 @@ func TestAuthorizeOAuthUser(t *testing.T) {
"token": model.NewId(), "token": model.NewId(),
}))) })))
_, _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, nil, nil, model.ServiceGitlab, "", state, "") _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, nil, nil, model.ServiceGitlab, "", state, "")
require.NotNil(t, err) require.NotNil(t, err)
assert.Equal(t, "api.oauth.invalid_state_token.app_error", err.Id) assert.Equal(t, "api.oauth.invalid_state_token.app_error", err.Id)
assert.Error(t, err.Unwrap()) assert.Error(t, err.Unwrap())
@@ -232,7 +232,7 @@ func TestAuthorizeOAuthUser(t *testing.T) {
state := makeState(token) state := makeState(token)
_, _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, nil, nil, model.ServiceGitlab, "", state, "") _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, nil, nil, model.ServiceGitlab, "", state, "")
require.NotNil(t, err) require.NotNil(t, err)
assert.Equal(t, "api.oauth.invalid_state_token.app_error", err.Id) assert.Equal(t, "api.oauth.invalid_state_token.app_error", err.Id)
assert.Equal(t, "", err.DetailedError) assert.Equal(t, "", err.DetailedError)
@@ -255,7 +255,7 @@ func TestAuthorizeOAuthUser(t *testing.T) {
"token": token.Token, "token": token.Token,
}))) })))
_, _, _, _, err = th.App.AuthorizeOAuthUser(th.Context, nil, nil, model.ServiceGitlab, "", state, "") _, _, _, err = th.App.AuthorizeOAuthUser(th.Context, nil, nil, model.ServiceGitlab, "", state, "")
require.NotNil(t, err) require.NotNil(t, err)
assert.Equal(t, "api.user.authorize_oauth_user.invalid_state.app_error", err.Id) assert.Equal(t, "api.user.authorize_oauth_user.invalid_state.app_error", err.Id)
}) })
@@ -268,7 +268,7 @@ func TestAuthorizeOAuthUser(t *testing.T) {
request := makeRequest("") request := makeRequest("")
state := makeState(makeToken(th, cookie)) state := makeState(makeToken(th, cookie))
_, _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, nil, request, model.ServiceGitlab, "", state, "") _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, nil, request, model.ServiceGitlab, "", state, "")
require.NotNil(t, err) require.NotNil(t, err)
assert.Equal(t, "api.user.authorize_oauth_user.invalid_state.app_error", err.Id) assert.Equal(t, "api.user.authorize_oauth_user.invalid_state.app_error", err.Id)
}) })
@@ -285,7 +285,7 @@ func TestAuthorizeOAuthUser(t *testing.T) {
request := makeRequest(cookie) request := makeRequest(cookie)
state := makeState(token) state := makeState(token)
_, _, _, _, err = th.App.AuthorizeOAuthUser(th.Context, nil, request, model.ServiceGitlab, "", state, "") _, _, _, err = th.App.AuthorizeOAuthUser(th.Context, nil, request, model.ServiceGitlab, "", state, "")
require.NotNil(t, err) require.NotNil(t, err)
assert.Equal(t, "api.user.authorize_oauth_user.invalid_state.app_error", err.Id) assert.Equal(t, "api.user.authorize_oauth_user.invalid_state.app_error", err.Id)
}) })
@@ -298,7 +298,7 @@ func TestAuthorizeOAuthUser(t *testing.T) {
request := makeRequest(cookie) request := makeRequest(cookie)
state := makeState(makeToken(th, cookie)) state := makeState(makeToken(th, cookie))
_, _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, &httptest.ResponseRecorder{}, request, model.ServiceGitlab, "", state, "") _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, &httptest.ResponseRecorder{}, request, model.ServiceGitlab, "", state, "")
require.NotNil(t, err) require.NotNil(t, err)
assert.Equal(t, "api.user.authorize_oauth_user.token_failed.app_error", err.Id) assert.Equal(t, "api.user.authorize_oauth_user.token_failed.app_error", err.Id)
}) })
@@ -316,7 +316,7 @@ func TestAuthorizeOAuthUser(t *testing.T) {
request := makeRequest(cookie) request := makeRequest(cookie)
state := makeState(makeToken(th, cookie)) state := makeState(makeToken(th, cookie))
_, _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, &httptest.ResponseRecorder{}, request, model.ServiceGitlab, "", state, "") _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, &httptest.ResponseRecorder{}, request, model.ServiceGitlab, "", state, "")
require.NotNil(t, err) require.NotNil(t, err)
assert.Equal(t, "api.user.authorize_oauth_user.bad_response.app_error", err.Id) assert.Equal(t, "api.user.authorize_oauth_user.bad_response.app_error", err.Id)
assert.Contains(t, err.DetailedError, "status_code=418") assert.Contains(t, err.DetailedError, "status_code=418")
@@ -336,7 +336,7 @@ func TestAuthorizeOAuthUser(t *testing.T) {
request := makeRequest(cookie) request := makeRequest(cookie)
state := makeState(makeToken(th, cookie)) state := makeState(makeToken(th, cookie))
_, _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, &httptest.ResponseRecorder{}, request, model.ServiceGitlab, "", state, "") _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, &httptest.ResponseRecorder{}, request, model.ServiceGitlab, "", state, "")
require.NotNil(t, err) require.NotNil(t, err)
assert.Equal(t, "api.user.authorize_oauth_user.bad_response.app_error", err.Id) assert.Equal(t, "api.user.authorize_oauth_user.bad_response.app_error", err.Id)
assert.Contains(t, err.DetailedError, "response_body=invalid") assert.Contains(t, err.DetailedError, "response_body=invalid")
@@ -359,7 +359,7 @@ func TestAuthorizeOAuthUser(t *testing.T) {
request := makeRequest(cookie) request := makeRequest(cookie)
state := makeState(makeToken(th, cookie)) state := makeState(makeToken(th, cookie))
_, _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, &httptest.ResponseRecorder{}, request, model.ServiceGitlab, "", state, "") _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, &httptest.ResponseRecorder{}, request, model.ServiceGitlab, "", state, "")
require.NotNil(t, err) require.NotNil(t, err)
assert.Equal(t, "api.user.authorize_oauth_user.bad_token.app_error", err.Id) assert.Equal(t, "api.user.authorize_oauth_user.bad_token.app_error", err.Id)
}) })
@@ -381,7 +381,7 @@ func TestAuthorizeOAuthUser(t *testing.T) {
request := makeRequest(cookie) request := makeRequest(cookie)
state := makeState(makeToken(th, cookie)) state := makeState(makeToken(th, cookie))
_, _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, &httptest.ResponseRecorder{}, request, model.ServiceGitlab, "", state, "") _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, &httptest.ResponseRecorder{}, request, model.ServiceGitlab, "", state, "")
require.NotNil(t, err) require.NotNil(t, err)
assert.Equal(t, "api.user.authorize_oauth_user.missing.app_error", err.Id) assert.Equal(t, "api.user.authorize_oauth_user.missing.app_error", err.Id)
}) })
@@ -403,7 +403,7 @@ func TestAuthorizeOAuthUser(t *testing.T) {
request := makeRequest(cookie) request := makeRequest(cookie)
state := makeState(makeToken(th, cookie)) state := makeState(makeToken(th, cookie))
_, _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, &httptest.ResponseRecorder{}, request, model.ServiceGitlab, "", state, "") _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, &httptest.ResponseRecorder{}, request, model.ServiceGitlab, "", state, "")
require.NotNil(t, err) require.NotNil(t, err)
assert.Equal(t, "api.user.authorize_oauth_user.service.app_error", err.Id) assert.Equal(t, "api.user.authorize_oauth_user.service.app_error", err.Id)
}) })
@@ -432,7 +432,7 @@ func TestAuthorizeOAuthUser(t *testing.T) {
request := makeRequest(cookie) request := makeRequest(cookie)
state := makeState(makeToken(th, cookie)) state := makeState(makeToken(th, cookie))
_, _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, &httptest.ResponseRecorder{}, request, model.ServiceGitlab, "", state, "") _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, &httptest.ResponseRecorder{}, request, model.ServiceGitlab, "", state, "")
require.NotNil(t, err) require.NotNil(t, err)
assert.Equal(t, "api.user.authorize_oauth_user.response.app_error", err.Id) assert.Equal(t, "api.user.authorize_oauth_user.response.app_error", err.Id)
}) })
@@ -463,7 +463,7 @@ func TestAuthorizeOAuthUser(t *testing.T) {
request := makeRequest(cookie) request := makeRequest(cookie)
state := makeState(makeToken(th, cookie)) state := makeState(makeToken(th, cookie))
_, _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, &httptest.ResponseRecorder{}, request, model.ServiceGitlab, "", state, "") _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, &httptest.ResponseRecorder{}, request, model.ServiceGitlab, "", state, "")
require.NotNil(t, err) require.NotNil(t, err)
assert.Equal(t, "oauth.gitlab.tos.error", err.Id) assert.Equal(t, "oauth.gitlab.tos.error", err.Id)
}) })
@@ -480,7 +480,7 @@ func TestAuthorizeOAuthUser(t *testing.T) {
providerMock.On("GetSSOSettings", mock.AnythingOfType("*request.Context"), mock.Anything, model.ServiceOpenid).Return(nil, errors.New("error")) providerMock.On("GetSSOSettings", mock.AnythingOfType("*request.Context"), mock.Anything, model.ServiceOpenid).Return(nil, errors.New("error"))
einterfaces.RegisterOAuthProvider(model.ServiceOpenid, providerMock) einterfaces.RegisterOAuthProvider(model.ServiceOpenid, providerMock)
_, _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, nil, nil, model.ServiceOpenid, "", "", "") _, _, _, err := th.App.AuthorizeOAuthUser(th.Context, nil, nil, model.ServiceOpenid, "", "", "")
require.NotNil(t, err) require.NotNil(t, err)
assert.Equal(t, "api.user.get_authorization_code.endpoint.app_error", err.Id) assert.Equal(t, "api.user.get_authorization_code.endpoint.app_error", err.Id)
}) })
@@ -532,14 +532,14 @@ func TestAuthorizeOAuthUser(t *testing.T) {
state := base64.StdEncoding.EncodeToString([]byte(model.MapToJSON(stateProps))) state := base64.StdEncoding.EncodeToString([]byte(model.MapToJSON(stateProps)))
recorder := httptest.ResponseRecorder{} recorder := httptest.ResponseRecorder{}
body, receivedTeamID, receivedStateProps, _, err := th.App.AuthorizeOAuthUser(th.Context, &recorder, request, model.ServiceGitlab, "", state, "") body, receivedStateProps, _, err := th.App.AuthorizeOAuthUser(th.Context, &recorder, request, model.ServiceGitlab, "", state, "")
require.NotNil(t, body) require.NotNil(t, body)
bodyBytes, bodyErr := io.ReadAll(body) bodyBytes, bodyErr := io.ReadAll(body)
require.NoError(t, bodyErr) require.NoError(t, bodyErr)
assert.Equal(t, userData, string(bodyBytes)) assert.Equal(t, userData, string(bodyBytes))
assert.Equal(t, stateProps["team_id"], receivedTeamID) // team_id is no longer returned as it was removed for security reasons
assert.Equal(t, stateProps, receivedStateProps) assert.Equal(t, stateProps, receivedStateProps)
assert.Nil(t, err) assert.Nil(t, err)

Просмотреть файл

@@ -720,6 +720,10 @@ func (a *App) AddUserToTeamByInviteId(c request.CTX, inviteId string, userID str
} }
team := teamChanResult.Data team := teamChanResult.Data
if team.IsGroupConstrained() {
return nil, nil, model.NewAppError("AddUserToTeamByInviteId", "app.team.invite_id.group_constrained.error", nil, "", http.StatusForbidden)
}
userChanResult := <-uchan userChanResult := <-uchan
if userChanResult.NErr != nil { if userChanResult.NErr != nil {
var nfErr *store.ErrNotFound var nfErr *store.ErrNotFound
@@ -1084,15 +1088,11 @@ func (a *App) AddTeamMemberByToken(c request.CTX, userID, tokenID string) (*mode
} }
func (a *App) AddTeamMemberByInviteId(c request.CTX, inviteId, userID string) (*model.TeamMember, *model.AppError) { func (a *App) AddTeamMemberByInviteId(c request.CTX, inviteId, userID string) (*model.TeamMember, *model.AppError) {
team, teamMember, err := a.AddUserToTeamByInviteId(c, inviteId, userID) _, teamMember, err := a.AddUserToTeamByInviteId(c, inviteId, userID)
if err != nil { if err != nil {
return nil, err return nil, err
} }
if team.IsGroupConstrained() {
return nil, model.NewAppError("AddTeamMemberByInviteId", "app.team.invite_id.group_constrained.error", nil, "", http.StatusForbidden)
}
return teamMember, nil return teamMember, nil
} }

Просмотреть файл

@@ -348,7 +348,7 @@ func (a *App) createUserOrGuest(c request.CTX, user *model.User, guest bool) (*m
return ruser, nil return ruser, nil
} }
func (a *App) CreateOAuthUser(c request.CTX, service string, userData io.Reader, teamID string, tokenUser *model.User) (*model.User, *model.AppError) { func (a *App) CreateOAuthUser(c request.CTX, service string, userData io.Reader, inviteToken string, inviteId string, tokenUser *model.User) (*model.User, *model.AppError) {
if !*a.Config().TeamSettings.EnableUserCreation { if !*a.Config().TeamSettings.EnableUserCreation {
return nil, model.NewAppError("CreateOAuthUser", "api.user.create_user.disabled.app_error", nil, "", http.StatusNotImplemented) return nil, model.NewAppError("CreateOAuthUser", "api.user.create_user.disabled.app_error", nil, "", http.StatusNotImplemented)
} }
@@ -401,19 +401,40 @@ func (a *App) CreateOAuthUser(c request.CTX, service string, userData io.Reader,
return nil, err return nil, err
} }
if teamID != "" { if err = a.AddUserToTeamByInviteIfNeeded(c, ruser, inviteToken, inviteId); err != nil {
err = a.AddUserToTeamByTeamId(c, teamID, user) c.Logger().Warn("Failed to add user to team", mlog.Err(err))
if err != nil { }
return nil, err
}
err = a.AddDirectChannels(c, teamID, user) return ruser, nil
}
func (a *App) AddUserToTeamByInviteIfNeeded(c request.CTX, user *model.User, inviteToken string, inviteId string) *model.AppError {
var team *model.Team
var err *model.AppError
if inviteToken != "" {
team, _, err = a.AddUserToTeamByToken(c, user.Id, inviteToken)
if err != nil { if err != nil {
c.Logger().Warn("Failed to add user to team using invite token", mlog.Err(err))
return err
}
} else if inviteId != "" {
team, _, err = a.AddUserToTeamByInviteId(c, inviteId, user.Id)
if err != nil {
c.Logger().Warn("Failed to add user to team using invite ID", mlog.Err(err))
return err
}
} else {
return nil
}
if team != nil {
if err = a.AddDirectChannels(c, team.Id, user); err != nil {
c.Logger().Warn("Failed to add direct channels", mlog.Err(err)) c.Logger().Warn("Failed to add direct channels", mlog.Err(err))
} }
} }
return ruser, nil return nil
} }
func (a *App) GetUser(userID string) (*model.User, *model.AppError) { func (a *App) GetUser(userID string) (*model.User, *model.AppError) {

Просмотреть файл

@@ -44,7 +44,7 @@ func TestCreateOAuthUser(t *testing.T) {
js, jsonErr := json.Marshal(glUser) js, jsonErr := json.Marshal(glUser)
require.NoError(t, jsonErr) require.NoError(t, jsonErr)
user, err := th.App.CreateOAuthUser(th.Context, model.UserAuthServiceGitlab, bytes.NewReader(js), th.BasicTeam.Id, nil) user, err := th.App.CreateOAuthUser(th.Context, model.UserAuthServiceGitlab, bytes.NewReader(js), "", "", nil)
require.Nil(t, err) require.Nil(t, err)
require.Equal(t, glUser.Username, user.Username, "usernames didn't match") require.Equal(t, glUser.Username, user.Username, "usernames didn't match")
@@ -73,7 +73,7 @@ func TestCreateOAuthUser(t *testing.T) {
assert.Equal(t, dbUser.Id, s) assert.Equal(t, dbUser.Id, s)
// data passed doesn't matter as return is mocked // data passed doesn't matter as return is mocked
_, err := th.App.CreateOAuthUser(th.Context, model.ServiceOffice365, strings.NewReader("{}"), th.BasicTeam.Id, nil) _, err := th.App.CreateOAuthUser(th.Context, model.ServiceOffice365, strings.NewReader("{}"), "", "", nil)
assert.Nil(t, err) assert.Nil(t, err)
u, er := th.App.Srv().Store().User().GetByEmail(dbUser.Email) u, er := th.App.Srv().Store().User().GetByEmail(dbUser.Email)
assert.NoError(t, er) assert.NoError(t, er)
@@ -83,7 +83,7 @@ func TestCreateOAuthUser(t *testing.T) {
t.Run("user creation disabled", func(t *testing.T) { t.Run("user creation disabled", func(t *testing.T) {
*th.App.Config().TeamSettings.EnableUserCreation = false *th.App.Config().TeamSettings.EnableUserCreation = false
_, err := th.App.CreateOAuthUser(th.Context, model.UserAuthServiceGitlab, strings.NewReader("{}"), th.BasicTeam.Id, nil) _, err := th.App.CreateOAuthUser(th.Context, model.UserAuthServiceGitlab, strings.NewReader("{}"), "", "", nil)
require.NotNil(t, err, "should have failed - user creation disabled") require.NotNil(t, err, "should have failed - user creation disabled")
}) })
} }
@@ -745,7 +745,7 @@ func createGitlabUser(t *testing.T, a *App, c request.CTX, id int64, username st
var user *model.User var user *model.User
var err *model.AppError var err *model.AppError
user, err = a.CreateOAuthUser(c, "gitlab", bytes.NewReader(gitlabUser), "", nil) user, err = a.CreateOAuthUser(c, "gitlab", bytes.NewReader(gitlabUser), "", "", nil)
require.Nil(t, err, "unable to create the user", err) require.Nil(t, err, "unable to create the user", err)
return user, gitlabUserObj return user, gitlabUserObj

Просмотреть файл

@@ -306,7 +306,7 @@ func completeOAuth(c *Context, w http.ResponseWriter, r *http.Request) {
uri := c.GetSiteURLHeader() + "/signup/" + service + "/complete" uri := c.GetSiteURLHeader() + "/signup/" + service + "/complete"
body, teamId, props, tokenUser, err := c.App.AuthorizeOAuthUser(c.AppContext, w, r, service, code, state, uri) body, props, tokenUser, err := c.App.AuthorizeOAuthUser(c.AppContext, w, r, service, code, state, uri)
action := "" action := ""
hasRedirectURL := false hasRedirectURL := false
@@ -337,7 +337,7 @@ func completeOAuth(c *Context, w http.ResponseWriter, r *http.Request) {
return return
} }
user, err := c.App.CompleteOAuth(c.AppContext, service, body, teamId, props, tokenUser) user, err := c.App.CompleteOAuth(c.AppContext, service, body, props, tokenUser)
if err != nil { if err != nil {
err.Translate(c.AppContext.T) err.Translate(c.AppContext.T)
c.LogErrorByCode(err) c.LogErrorByCode(err)
@@ -449,13 +449,11 @@ func loginWithOAuth(c *Context, w http.ResponseWriter, r *http.Request) {
auditRec.AddMeta("service", c.Params.Service) auditRec.AddMeta("service", c.Params.Service)
defer c.LogAuditRec(auditRec) defer c.LogAuditRec(auditRec)
teamId, err := c.App.GetTeamIdFromQuery(c.AppContext, r.URL.Query()) // Get invite token or ID instead of team_id
if err != nil { tokenID := r.URL.Query().Get("t")
c.Err = err inviteId := r.URL.Query().Get("id")
return
}
authURL, err := c.App.GetOAuthLoginEndpoint(c.AppContext, w, r, c.Params.Service, teamId, model.OAuthActionLogin, redirectURL, loginHint, false, desktopToken) authURL, err := c.App.GetOAuthLoginEndpoint(c.AppContext, w, r, c.Params.Service, model.OAuthActionLogin, redirectURL, loginHint, false, desktopToken, tokenID, inviteId)
if err != nil { if err != nil {
c.Err = err c.Err = err
return return
@@ -485,13 +483,11 @@ func mobileLoginWithOAuth(c *Context, w http.ResponseWriter, r *http.Request) {
auditRec.AddMeta("service", c.Params.Service) auditRec.AddMeta("service", c.Params.Service)
defer c.LogAuditRec(auditRec) defer c.LogAuditRec(auditRec)
teamId, err := c.App.GetTeamIdFromQuery(c.AppContext, r.URL.Query()) // Get invite token or ID instead of team_id
if err != nil { tokenID := r.URL.Query().Get("t")
c.Err = err inviteId := r.URL.Query().Get("id")
return
}
authURL, err := c.App.GetOAuthLoginEndpoint(c.AppContext, w, r, c.Params.Service, teamId, model.OAuthActionMobile, redirectURL, "", true, "") authURL, err := c.App.GetOAuthLoginEndpoint(c.AppContext, w, r, c.Params.Service, model.OAuthActionMobile, redirectURL, "", true, "", tokenID, inviteId)
if err != nil { if err != nil {
c.Err = err c.Err = err
return return
@@ -520,15 +516,13 @@ func signupWithOAuth(c *Context, w http.ResponseWriter, r *http.Request) {
auditRec.AddMeta("service", c.Params.Service) auditRec.AddMeta("service", c.Params.Service)
defer c.LogAuditRec(auditRec) defer c.LogAuditRec(auditRec)
teamId, err := c.App.GetTeamIdFromQuery(c.AppContext, r.URL.Query()) // Get invite token or ID instead of team_id
if err != nil { tokenID := r.URL.Query().Get("t")
c.Err = err inviteId := r.URL.Query().Get("id")
return
}
desktopToken := r.URL.Query().Get("desktop_token") desktopToken := r.URL.Query().Get("desktop_token")
authURL, err := c.App.GetOAuthSignupEndpoint(c.AppContext, w, r, c.Params.Service, teamId, desktopToken) authURL, err := c.App.GetOAuthSignupEndpoint(c.AppContext, w, r, c.Params.Service, desktopToken, tokenID, inviteId)
if err != nil { if err != nil {
c.Err = err c.Err = err
return return

Просмотреть файл

@@ -31,11 +31,9 @@ func loginWithSaml(c *Context, w http.ResponseWriter, r *http.Request) {
return return
} }
teamId, err := c.App.GetTeamIdFromQuery(c.AppContext, r.URL.Query()) tokenID := r.URL.Query().Get("t")
if err != nil { inviteId := r.URL.Query().Get("id")
c.Err = err
return
}
action := r.URL.Query().Get("action") action := r.URL.Query().Get("action")
isMobile := action == model.OAuthActionMobile isMobile := action == model.OAuthActionMobile
redirectURL := html.EscapeString(r.URL.Query().Get("redirect_to")) redirectURL := html.EscapeString(r.URL.Query().Get("redirect_to"))
@@ -43,7 +41,11 @@ func loginWithSaml(c *Context, w http.ResponseWriter, r *http.Request) {
relayState := "" relayState := ""
if action != "" { if action != "" {
relayProps["team_id"] = teamId if tokenID != "" {
relayProps["invite_token"] = tokenID
} else if inviteId != "" {
relayProps["invite_id"] = inviteId
}
relayProps["action"] = action relayProps["action"] = action
if action == model.OAuthActionEmailToSSO { if action == model.OAuthActionEmailToSSO {
relayProps["email_token"] = r.URL.Query().Get("email_token") relayProps["email_token"] = r.URL.Query().Get("email_token")
@@ -149,14 +151,11 @@ func completeSaml(c *Context, w http.ResponseWriter, r *http.Request) {
switch action { switch action {
case model.OAuthActionSignup: case model.OAuthActionSignup:
if teamId := relayProps["team_id"]; teamId != "" { inviteToken := relayProps["invite_token"]
if err = c.App.AddUserToTeamByTeamId(c.AppContext, teamId, user); err != nil { inviteId := relayProps["invite_id"]
c.LogErrorByCode(err) if err = c.App.AddUserToTeamByInviteIfNeeded(c.AppContext, user, inviteToken, inviteId); err != nil {
break c.LogErrorByCode(err)
} break
if err = c.App.AddDirectChannels(c.AppContext, teamId, user); err != nil {
c.LogErrorByCode(err)
}
} }
case model.OAuthActionEmailToSSO: case model.OAuthActionEmailToSSO:
if err = c.App.RevokeAllSessions(c.AppContext, user.Id); err != nil { if err = c.App.RevokeAllSessions(c.AppContext, user.Id); err != nil {