From af6e968980e561a1d490a9a49798a203b68cc3ce Mon Sep 17 00:00:00 2001 From: Qiu Jian Date: Tue, 22 Feb 2022 02:36:40 +0800 Subject: [PATCH] fix: allow IDP initiated SAML login --- cmd/climc/shell/identity/identityproviders.go | 14 +++ pkg/apigateway/handler/auth.go | 8 +- pkg/apigateway/handler/idp.go | 106 +++++++++++++----- pkg/apis/identity/identityprovider.go | 12 ++ pkg/apis/identity/saml.go | 9 ++ pkg/keystone/driver/base.go | 4 + pkg/keystone/driver/driver.go | 1 + pkg/keystone/driver/saml/class.go | 1 + pkg/keystone/driver/saml/saml.go | 8 ++ pkg/keystone/models/identity_provider.go | 23 ++++ 10 files changed, 156 insertions(+), 30 deletions(-) diff --git a/cmd/climc/shell/identity/identityproviders.go b/cmd/climc/shell/identity/identityproviders.go index c39b114376..dd8331c46a 100644 --- a/cmd/climc/shell/identity/identityproviders.go +++ b/cmd/climc/shell/identity/identityproviders.go @@ -748,6 +748,20 @@ func init() { return nil }) + type IdpGetCallbackUriOptions struct { + ID string `help:"id or name of idp to query" json:"-"` + + api.GetIdpSsoCallbackUriInput + } + R(&IdpGetCallbackUriOptions{}, "idp-sso-callback-url", "Get sso callback url of a SSO idp", func(s *mcclient.ClientSession, args *IdpGetCallbackUriOptions) error { + result, err := modules.IdentityProviders.GetSpecific(s, args.ID, "sso-callback-uri", jsonutils.Marshal(args)) + if err != nil { + return err + } + printObject(result) + return nil + }) + type IdpSetDefaultSsoOptions struct { ID string `help:"id or name of idp to set default Sso" json:"-"` diff --git a/pkg/apigateway/handler/auth.go b/pkg/apigateway/handler/auth.go index d4e0b612c9..8980bdf1e4 100644 --- a/pkg/apigateway/handler/auth.go +++ b/pkg/apigateway/handler/auth.go @@ -69,6 +69,7 @@ func (h *AuthHandlers) AddMethods() { NewHP(h.getIdpSsoRedirectUri, "sso", "redirect", ""), NewHP(h.listTotpRecoveryQuestions, "recovery"), NewHP(h.handleSsoLogin, "ssologin"), + NewHP(h.handleIdpInitSsoLogin, "ssologin", ""), NewHP(h.postLogoutHandler, "logout"), // oidc auth NewHP(handleOIDCAuth, "oidc", "auth"), @@ -84,6 +85,7 @@ func (h *AuthHandlers) AddMethods() { NewHP(h.postLoginHandler, "login"), NewHP(h.postLogoutHandler, "logout"), NewHP(h.handleSsoLogin, "ssologin"), + NewHP(h.handleIdpInitSsoLogin, "ssologin", ""), NewHP(handleOIDCToken, "oidc", "token"), ) @@ -347,7 +349,9 @@ func (h *AuthHandlers) doCredentialLogin(ctx context.Context, req *http.Request, domain, _ := body.GetString("domain") token, err = auth.Client().AuthenticateWeb(uname, passwd, domain, "", "", cliIp) } else if body.Contains("idp_driver") { // sso login - token, err = processSsoLoginData(body, cliIp) + idpId, _ := body.GetString("idp_id") + redirectUri := getSsoCallbackUrl(ctx, req, idpId) + token, err = processSsoLoginData(body, cliIp, redirectUri) if err != nil { return nil, errors.Wrap(err, "processSsoLoginData") } @@ -1052,7 +1056,7 @@ func getUserInfo2(s *mcclient.ClientSession, uid string, pid string, loginIp str data.Add(jsonutils.JSONFalse, "enable_quota_check") } - data.Add(jsonutils.NewString(getSsoCallbackUrl()), "sso_callback_url") + // data.Add(jsonutils.NewString(getSsoCallbackUrl(ctx, req, idpId)), "sso_callback_url") return data, nil } diff --git a/pkg/apigateway/handler/idp.go b/pkg/apigateway/handler/idp.go index 00b1c577c9..9675adbb7a 100644 --- a/pkg/apigateway/handler/idp.go +++ b/pkg/apigateway/handler/idp.go @@ -40,13 +40,30 @@ import ( "yunion.io/x/onecloud/pkg/util/netutils2" ) -func getSsoCallbackUrl() string { +func getSsoBaseCallbackUrl() string { if options.Options.SsoRedirectUrl == "" { return httputils.JoinPath(options.Options.ApiServer, "api/v1/auth/ssologin") } return options.Options.SsoRedirectUrl } +func getSsoCallbackUrl(ctx context.Context, req *http.Request, idpId string) string { + baseUrl := getSsoBaseCallbackUrl() + s := auth.GetAdminSession(ctx, FetchRegion(req), "") + input := api.GetIdpSsoCallbackUriInput{ + RedirectUri: baseUrl, + } + resp, err := modules.IdentityProviders.GetSpecific(s, idpId, "sso-callback-uri", jsonutils.Marshal(input)) + if err != nil { + return baseUrl + } + ret, err := resp.GetString("redirect_uri") + if err != nil { + return baseUrl + } + return ret +} + func getSsoAuthCallbackUrl() string { if options.Options.SsoAuthCallbackUrl == "" { return httputils.JoinPath(options.Options.ApiServer, "auth") @@ -92,7 +109,7 @@ func (h *AuthHandlers) getIdpSsoRedirectUri(ctx context.Context, w http.Response } query.(*jsonutils.JSONDict).Set("idp_nonce", jsonutils.NewString(utils.GenRequestId(4))) state := base64.URLEncoding.EncodeToString([]byte(query.String())) - redirectUri := getSsoCallbackUrl() + redirectUri := getSsoCallbackUrl(ctx, req, idpId) s := auth.GetAdminSession(ctx, FetchRegion(req), "") input := api.GetIdpSsoRedirectUriInput{ RedirectUri: redirectUri, @@ -125,13 +142,38 @@ func findExtUserId(input string) string { return "" } +func (h *AuthHandlers) handleIdpInitSsoLogin(ctx context.Context, w http.ResponseWriter, req *http.Request) { + params := appctx.AppContextParams(ctx) + idpId := params[""] + s := auth.GetAdminSession(ctx, FetchRegion(req), "") + resp, err := modules.IdentityProviders.Get(s, idpId, nil) + if err != nil { + httperrors.GeneralServerError(ctx, w, err) + return + } + idpDriver, _ := resp.GetString("driver") + h.internalSsoLogin(ctx, w, req, idpId, idpDriver) +} + func (h *AuthHandlers) handleSsoLogin(ctx context.Context, w http.ResponseWriter, req *http.Request) { - idpId := getCookie(req, "idp_id") - idpDriver := getCookie(req, "idp_driver") + + h.internalSsoLogin(ctx, w, req, "", "") +} + +func (h *AuthHandlers) internalSsoLogin(ctx context.Context, w http.ResponseWriter, req *http.Request, idpId, idpDriver string) { + idpIdC := getCookie(req, "idp_id") + idpDriverC := getCookie(req, "idp_driver") idpState := getCookie(req, "idp_state") idpReferer := getCookie(req, "idp_referer") idpLinkUser := getCookie(req, "idp_link_user") + if len(idpIdC) > 0 { + idpId = idpIdC + } + if len(idpDriverC) > 0 { + idpDriver = idpDriverC + } + for _, k := range []string{"idp_id", "idp_driver", "idp_state", "idp_referer", "idp_link_user"} { clearCookie(w, k, "") } @@ -143,20 +185,23 @@ func (h *AuthHandlers) handleSsoLogin(ctx context.Context, w http.ResponseWriter if len(idpDriver) == 0 { missing = append(missing, "idp_driver") } - if len(idpState) == 0 { + /*if len(idpState) == 0 { missing = append(missing, "idp_state") } if len(idpReferer) == 0 { missing = append(missing, "idp_referer") - } + }*/ if len(missing) > 0 { httperrors.TimeoutError(ctx, w, "session expires, missing %s", strings.Join(missing, ",")) return } - idpStateQsBytes, _ := base64.URLEncoding.DecodeString(idpState) - idpStateQs, _ := jsonutils.Parse(idpStateQsBytes) - log.Debugf("state query sting: %s", idpStateQs) + var idpStateQs jsonutils.JSONObject + if len(idpState) > 0 { + idpStateQsBytes, _ := base64.URLEncoding.DecodeString(idpState) + idpStateQs, _ = jsonutils.Parse(idpStateQsBytes) + log.Debugf("state query sting: %s", idpStateQs) + } var body jsonutils.JSONObject var err error @@ -212,11 +257,12 @@ func (h *AuthHandlers) handleSsoLogin(ctx context.Context, w http.ResponseWriter } } refererUrl, _ := url.Parse(referer) - if refererUrl == nil { + if refererUrl == nil && len(idpReferer) > 0 { refererUrl, _ = url.Parse(idpReferer) } - if err != nil { - log.Debugf("error: %s refererUrl: %s", err, refererUrl) + if refererUrl == nil { + httperrors.InvalidInputError(ctx, w, "empty referer link") + return } redirUrl := generateRedirectUrl(refererUrl, idpStateQs, err, idpId, idpUserId) appsrv.SendRedirect(w, redirUrl) @@ -258,7 +304,7 @@ func generateRedirectUrl(originUrl *url.URL, stateQs jsonutils.JSONObject, err e return originUrl.String() } -func processSsoLoginData(body jsonutils.JSONObject, cliIp string) (mcclient.TokenCredential, error) { +func processSsoLoginData(body jsonutils.JSONObject, cliIp string, redirectUri string) (mcclient.TokenCredential, error) { var token mcclient.TokenCredential var err error idpDriver, _ := body.GetString("idp_driver") @@ -266,7 +312,6 @@ func processSsoLoginData(body jsonutils.JSONObject, cliIp string) (mcclient.Toke idpState, _ := body.GetString("idp_state") switch idpDriver { case api.IdentityDriverCAS: - redirectUri := getSsoCallbackUrl() ticket, _ := body.GetString("ticket") if len(ticket) == 0 { return nil, httperrors.NewMissingParameterError("ticket") @@ -283,7 +328,6 @@ func processSsoLoginData(body jsonutils.JSONObject, cliIp string) (mcclient.Toke } token, err = auth.Client().AuthenticateSAML(idpId, samlResp, "", "", "", cliIp) case api.IdentityDriverOIDC: - redirectUri := getSsoCallbackUrl() code, _ := body.GetString("code") state, _ := body.GetString("state") if state != idpState { @@ -321,7 +365,8 @@ func linkWithExistingUser(ctx context.Context, req *http.Request, idpId, idpLink return errors.Wrap(httperrors.ErrConflict, "link user id inconsistent with credential") } cliIp := netutils2.GetHttpRequestIp(req) - ntoken, err := processSsoLoginData(body, cliIp) + redirectUri := getSsoCallbackUrl(ctx, req, idpId) + ntoken, err := processSsoLoginData(body, cliIp, redirectUri) if err != nil { if errors.Cause(err) != httperrors.ErrUserNotFound { return errors.Wrap(err, "invalid ssologin result") @@ -382,10 +427,9 @@ func handleUnlinkIdp(ctx context.Context, w http.ResponseWriter, req *http.Reque } func fetchIdpBasicConfig(ctx context.Context, w http.ResponseWriter, req *http.Request) { - s := auth.GetAdminSession(ctx, FetchRegion(req), "") params := appctx.AppContextParams(ctx) idpId := params[""] - info, err := getIdpBasicConfig(s, idpId) + info, err := getIdpBasicConfig(ctx, req, idpId) if err != nil { httperrors.GeneralServerError(ctx, w, err) return @@ -393,25 +437,31 @@ func fetchIdpBasicConfig(ctx context.Context, w http.ResponseWriter, req *http.R appsrv.SendJSON(w, info) } -func getIdpBasicConfig(s *mcclient.ClientSession, idpId string) (jsonutils.JSONObject, error) { - idp, err := modules.IdentityProviders.Get(s, idpId, nil) - if err != nil { - return nil, errors.Wrap(err, "Fetch") +func getIdpBasicConfig(ctx context.Context, req *http.Request, idpId string) (jsonutils.JSONObject, error) { + s := auth.GetAdminSession(ctx, FetchRegion(req), "") + baseUrl := getSsoBaseCallbackUrl() + input := api.GetIdpSsoCallbackUriInput{ + RedirectUri: baseUrl, } + resp, err := modules.IdentityProviders.GetSpecific(s, idpId, "sso-callback-uri", jsonutils.Marshal(input)) + if err != nil { + return nil, errors.Wrap(err, "GetSpecific sso-callback-uri") + } + redir, _ := resp.GetString("redirect_uri") + idpDriver, _ := resp.GetString("driver") info := jsonutils.NewDict() - idpDriver, _ := idp.GetString("driver") switch idpDriver { case api.IdentityDriverSQL: case api.IdentityDriverLDAP: case api.IdentityDriverCAS: - info.Add(jsonutils.NewString(getSsoCallbackUrl()), "redirect_uri") + info.Add(jsonutils.NewString(redir), "redirect_uri") case api.IdentityDriverSAML: info.Add(jsonutils.NewString(options.Options.ApiServer), "entity_id") - info.Add(jsonutils.NewString(getSsoCallbackUrl()), "redirect_uri") + info.Add(jsonutils.NewString(redir), "redirect_uri") case api.IdentityDriverOIDC: - info.Add(jsonutils.NewString(getSsoCallbackUrl()), "redirect_uri") + info.Add(jsonutils.NewString(redir), "redirect_uri") case api.IdentityDriverOAuth2: - info.Add(jsonutils.NewString(getSsoCallbackUrl()), "redirect_uri") + info.Add(jsonutils.NewString(redir), "redirect_uri") default: } return info, nil @@ -422,7 +472,7 @@ func fetchIdpSAMLMetadata(ctx context.Context, w http.ResponseWriter, req *http. params := appctx.AppContextParams(ctx) idpId := params[""] query := jsonutils.NewDict() - query.Set("redirect_uri", jsonutils.NewString(getSsoCallbackUrl())) + query.Set("redirect_uri", jsonutils.NewString(getSsoCallbackUrl(ctx, req, idpId))) md, err := modules.IdentityProviders.GetSpecific(s, idpId, "saml-metadata", query) if err != nil { httperrors.GeneralServerError(ctx, w, err) diff --git a/pkg/apis/identity/identityprovider.go b/pkg/apis/identity/identityprovider.go index 5e6e4644c3..863780a76a 100644 --- a/pkg/apis/identity/identityprovider.go +++ b/pkg/apis/identity/identityprovider.go @@ -126,3 +126,15 @@ type GetIdpSsoRedirectUriOutput struct { type PerformDefaultSsoInput struct { Enable *bool `json:"enable" help:"enable default sso" negative:"disable"` } + +type GetIdpSsoCallbackUriInput struct { + // SSO回调地址 + RedirectUri string `json:"redirect_uri"` +} + +type GetIdpSsoCallbackUriOutput struct { + // SSO回调地址 + RedirectUri string `json:"redirect_uri"` + // Driver + Driver string `json:"driver"` +} diff --git a/pkg/apis/identity/saml.go b/pkg/apis/identity/saml.go index 669cd44596..f7135fb271 100644 --- a/pkg/apis/identity/saml.go +++ b/pkg/apis/identity/saml.go @@ -32,20 +32,29 @@ type SIdpAttributeOptions struct { DefaultRoleId string `json:"default_role_id"` } +type SSAMLIdpBaseConfigOptions struct { + AllowIdpInit *bool `json:"allow_idp_init"` +} + type SSAMLIdpConfigOptions struct { EntityId string `json:"entity_id"` RedirectSSOUrl string `json:"redirect_sso_url"` + SSAMLIdpBaseConfigOptions + SIdpAttributeOptions } type SSAMLTestIdpConfigOptions struct { // empty + SSAMLIdpBaseConfigOptions } type SSAMLAzureADConfigOptions struct { TenantId string `json:"tenant_id"` + SSAMLIdpBaseConfigOptions + SIdpAttributeOptions } diff --git a/pkg/keystone/driver/base.go b/pkg/keystone/driver/base.go index 80abc53e54..91315a7a9c 100644 --- a/pkg/keystone/driver/base.go +++ b/pkg/keystone/driver/base.go @@ -58,6 +58,10 @@ func (base *SBaseIdentityDriver) IIdentityBackend() IIdentityBackend { return base.GetVirtualObject().(IIdentityBackend) } +func (base *SBaseIdentityDriver) GetSsoCallbackUri(callbackUrl string) string { + return callbackUrl +} + func NewBaseIdentityDriver(idpId, idpName, template, targetDomainId string, conf api.TConfigs) (SBaseIdentityDriver, error) { drv := SBaseIdentityDriver{} drv.IdpId = idpId diff --git a/pkg/keystone/driver/driver.go b/pkg/keystone/driver/driver.go index e18db69a57..d5b71f57c2 100644 --- a/pkg/keystone/driver/driver.go +++ b/pkg/keystone/driver/driver.go @@ -35,6 +35,7 @@ type IIdentityBackendClass interface { type IIdentityBackend interface { Authenticate(ctx context.Context, identity mcclient.SAuthenticationIdentity) (*api.SUserExtended, error) GetSsoRedirectUri(ctx context.Context, callbackUrl, state string) (string, error) + GetSsoCallbackUri(callbackUrl string) string Sync(ctx context.Context) error Probe(ctx context.Context) error } diff --git a/pkg/keystone/driver/saml/class.go b/pkg/keystone/driver/saml/class.go index 0e6ea52262..ba15c9c3b0 100644 --- a/pkg/keystone/driver/saml/class.go +++ b/pkg/keystone/driver/saml/class.go @@ -130,6 +130,7 @@ func (self *SSAMLDriverClass) ValidateConfig(ctx context.Context, userCred mccli if err != nil { return tconf, errors.Wrap(err, "Unmarshal new config") } + nconf["allow_idp_init"] = jsonutils.JSONTrue tconf[api.IdentityDriverSAML] = nconf return tconf, nil } diff --git a/pkg/keystone/driver/saml/saml.go b/pkg/keystone/driver/saml/saml.go index 081814b2d4..bb89439f07 100644 --- a/pkg/keystone/driver/saml/saml.go +++ b/pkg/keystone/driver/saml/saml.go @@ -29,6 +29,7 @@ import ( "yunion.io/x/onecloud/pkg/keystone/models" "yunion.io/x/onecloud/pkg/keystone/saml" "yunion.io/x/onecloud/pkg/mcclient" + "yunion.io/x/onecloud/pkg/util/httputils" "yunion.io/x/onecloud/pkg/util/samlutils" "yunion.io/x/onecloud/pkg/util/samlutils/sp" ) @@ -79,6 +80,13 @@ func (self *SSAMLDriver) prepareConfig() error { return nil } +func (self *SSAMLDriver) GetSsoCallbackUri(callbackUrl string) string { + if self.samlConfig.AllowIdpInit != nil && *self.samlConfig.AllowIdpInit { + callbackUrl = httputils.JoinPath(callbackUrl, self.IdpId) + } + return callbackUrl +} + func (self *SSAMLDriver) GetSsoRedirectUri(ctx context.Context, callbackUrl, state string) (string, error) { spLoginFunc := func(ctx context.Context, idp *sp.SSAMLIdentityProvider) (sp.SSAMLSpInitiatedLoginRequest, error) { result := sp.SSAMLSpInitiatedLoginRequest{} diff --git a/pkg/keystone/models/identity_provider.go b/pkg/keystone/models/identity_provider.go index 78a3e7d56c..3636e969eb 100644 --- a/pkg/keystone/models/identity_provider.go +++ b/pkg/keystone/models/identity_provider.go @@ -1382,6 +1382,29 @@ func (idp *SIdentityProvider) GetDetailsSsoRedirectUri(ctx context.Context, user return output, nil } +func (idp *SIdentityProvider) GetDetailsSsoCallbackUri(ctx context.Context, userCred mcclient.TokenCredential, query api.GetIdpSsoCallbackUriInput) (api.GetIdpSsoCallbackUriOutput, error) { + output := api.GetIdpSsoCallbackUriOutput{} + conf, err := GetConfigs(idp, true, nil, nil) + if err != nil { + return output, errors.Wrap(err, "idp.GetConfig") + } + + backend, err := driver.GetDriver(idp.Driver, idp.Id, idp.Name, idp.Template, idp.TargetDomainId, conf) + if err != nil { + return output, errors.Wrap(err, "driver.GetDriver") + } + + uri := backend.GetSsoCallbackUri(query.RedirectUri) + if err != nil { + return output, errors.Wrap(err, "backend.GetSsoCallbackUri") + } + + output.RedirectUri = uri + output.Driver = idp.Driver + + return output, nil +} + func (idp *SIdentityProvider) SyncOrCreateDomainAndUser(ctx context.Context, extDomainId, extDomainName string, extUsrId, extUsrName string) (*SDomain, *SUser, error) { var ( domain *SDomain