diff --git a/lib/auth/github.go b/lib/auth/github.go index cb946058c84..d2ab8cf8abc 100644 --- a/lib/auth/github.go +++ b/lib/auth/github.go @@ -39,6 +39,9 @@ import ( "github.com/sirupsen/logrus" ) +// ErrGithubNoTeams results from a github user not beloging to any teams. +var ErrGithubNoTeams = trace.BadParameter("user does not belong to any teams configured in connector; the configuration may have typos.") + // CreateGithubAuthRequest creates a new request for Github OAuth2 flow func (a *Server) CreateGithubAuthRequest(ctx context.Context, req types.GithubAuthRequest) (*types.GithubAuthRequest, error) { _, client, err := a.getGithubConnectorAndClient(ctx, req) @@ -453,9 +456,7 @@ func (a *Server) calculateGithubUser(connector types.GithubConnector, claims *ty // Calculate logins, kubegroups, roles, and traits. p.roles, p.kubeGroups, p.kubeUsers = connector.MapClaims(*claims) if len(p.roles) == 0 { - return nil, trace.BadParameter( - "user %q does not belong to any teams configured in %q connector; the configuration may have typos.", - claims.Username, connector.GetName()) + return nil, trace.Wrap(ErrGithubNoTeams) } p.traits = map[string][]string{ constants.TraitLogins: {p.username}, diff --git a/lib/auth/github_test.go b/lib/auth/github_test.go index 237dd9273e0..603a6b0e2a4 100644 --- a/lib/auth/github_test.go +++ b/lib/auth/github_test.go @@ -257,3 +257,27 @@ func (m *mockedGithubManager) validateGithubAuthCallback(ctx context.Context, di return nil, trace.NotImplemented("mockValidateGithubAuthCallback not implemented") } + +func TestCalculateGithubUserNoTeams(t *testing.T) { + a := &Server{} + connector, err := types.NewGithubConnector("github", types.GithubConnectorSpecV3{ + TeamsToRoles: []types.TeamRolesMapping{ + { + Organization: "org1", + Team: "teamx", + Roles: []string{"role"}, + }, + }, + }) + require.NoError(t, err) + + _, err = a.calculateGithubUser(connector, &types.GithubClaims{ + Username: "octocat", + OrganizationToTeams: map[string][]string{ + "org1": {"team1", "team2"}, + "org2": {"team1"}, + }, + Teams: []string{"team1", "team2", "team1"}, + }, &types.GithubAuthRequest{}) + require.ErrorIs(t, err, ErrGithubNoTeams) +} diff --git a/lib/auth/oidc.go b/lib/auth/oidc.go index fb02cc4ce84..32bd9fc5bbe 100644 --- a/lib/auth/oidc.go +++ b/lib/auth/oidc.go @@ -43,6 +43,9 @@ import ( "github.com/gravitational/trace" ) +// ErrOIDCNoRoles results from not mapping any roles from OIDC claims. +var ErrOIDCNoRoles = trace.AccessDenied("No roles mapped from claims. The mappings may contain typos.") + // getOIDCConnectorAndClient returns the associated oidc connector // and client for the given oidc auth request. func (a *Server) getOIDCConnectorAndClient(ctx context.Context, request types.OIDCAuthRequest) (types.OIDCConnector, *oidc.Client, error) { @@ -612,7 +615,7 @@ func (a *Server) calculateOIDCUser(diagCtx *ssoDiagContext, connector types.OIDC Message: "No roles mapped for the user. The mappings may contain typos.", } } - return nil, trace.AccessDenied("No roles mapped from claims. The mappings may contain typos.") + return nil, trace.Wrap(ErrOIDCNoRoles) } // Pick smaller for role: session TTL from role or requested TTL. diff --git a/lib/auth/oidc_test.go b/lib/auth/oidc_test.go index 9fb5832f5b9..7f4d6a1477b 100644 --- a/lib/auth/oidc_test.go +++ b/lib/auth/oidc_test.go @@ -195,144 +195,174 @@ func TestUserInfoBadStatus(t *testing.T) { func TestSSODiagnostic(t *testing.T) { t.Parallel() - - ctx := context.Background() - s := setUpSuite(t) - // Create configurable IdP to use in tests. - idp := newFakeIDP(t, false /* tls */) - - // create role referenced in request. - role, err := types.NewRole("access", types.RoleSpecV5{ - Allow: types.RoleConditions{ - Logins: []string{"dummy"}, - }, - }) - require.NoError(t, err) - err = s.a.CreateRole(role) - require.NoError(t, err) - - // connector spec - spec := types.OIDCConnectorSpecV3{ - IssuerURL: idp.s.URL, - ClientID: "00000000000000000000000000000000", - ClientSecret: "0000000000000000000000000000000000000000000000000000000000000000", - Display: "Test", - Scope: []string{"groups"}, - ClaimsToRoles: []types.ClaimMapping{ - { - Claim: "groups", - Value: "idp-admin", - Roles: []string{"access"}, + tests := []struct { + name string + claimsToRoles []types.ClaimMapping + wantValidateErr error + }{ + { + name: "success", + claimsToRoles: []types.ClaimMapping{ + { + Claim: "groups", + Value: "idp-admin", + Roles: []string{"access"}, + }, }, }, - RedirectURLs: []string{"https://proxy.example.com/v1/webapi/oidc/callback"}, - } - - oidcRequest := types.OIDCAuthRequest{ - ConnectorID: "-sso-test-okta", - Type: constants.OIDC, - CertTTL: defaults.OIDCAuthRequestTTL, - SSOTestFlow: true, - ConnectorSpec: &spec, - } - - request, err := s.a.CreateOIDCAuthRequest(ctx, oidcRequest) - require.NoError(t, err) - require.NotNil(t, request) - - values := url.Values{ - "code": []string{"XXX-code"}, - "state": []string{request.StateToken}, - } - - // override getClaimsFun. - s.a.getClaimsFun = func(closeCtx context.Context, oidcClient *oidc.Client, connector types.OIDCConnector, code string) (jose.Claims, error) { - cc := map[string]interface{}{ - "email_verified": true, - "groups": []string{"everyone", "idp-admin", "idp-dev"}, - "email": "superuser@example.com", - "sub": "00001234abcd", - "exp": 1652091713.0, - } - return cc, nil - } - - resp, err := s.a.ValidateOIDCAuthCallback(ctx, values) - require.NoError(t, err) - require.NotNil(t, resp) - require.Equal(t, &OIDCAuthResponse{ - Username: "superuser@example.com", - Identity: types.ExternalIdentity{ - ConnectorID: "-sso-test-okta", - Username: "superuser@example.com", - }, - Req: *request, - }, resp) - - diagCtx := ssoDiagContext{} - - resp, err = s.a.validateOIDCAuthCallback(ctx, &diagCtx, values) - require.NoError(t, err) - require.NotNil(t, resp) - require.Equal(t, &OIDCAuthResponse{ - Username: "superuser@example.com", - Identity: types.ExternalIdentity{ - ConnectorID: "-sso-test-okta", - Username: "superuser@example.com", - }, - Req: *request, - }, resp) - require.Equal(t, types.SSODiagnosticInfo{ - TestFlow: true, - Success: true, - CreateUserParams: &types.CreateUserParams{ - ConnectorName: "-sso-test-okta", - Username: "superuser@example.com", - Logins: nil, - KubeGroups: nil, - KubeUsers: nil, - Roles: []string{"access"}, - Traits: map[string][]string{ - "email": {"superuser@example.com"}, - "groups": {"everyone", "idp-admin", "idp-dev"}, - "sub": {"00001234abcd"}, + { + name: "fail to map claims to roles", + claimsToRoles: []types.ClaimMapping{ + { + Claim: "groups", + Value: "nonexistant", + Roles: []string{"access"}, + }, }, - SessionTTL: 600000000000, + wantValidateErr: ErrOIDCNoRoles, }, - OIDCClaimsToRoles: []types.ClaimMapping{ - { - Claim: "groups", - Value: "idp-admin", - Roles: []string{"access"}, - }, - }, - OIDCClaimsToRolesWarnings: nil, - OIDCClaims: map[string]interface{}{ - "email_verified": true, - "groups": []string{"everyone", "idp-admin", "idp-dev"}, - "email": "superuser@example.com", - "sub": "00001234abcd", - "exp": 1652091713.0, - }, - OIDCIdentity: &types.OIDCIdentity{ - ID: "00001234abcd", - Name: "", - Email: "superuser@example.com", - ExpiresAt: diagCtx.info.OIDCIdentity.ExpiresAt, - }, - OIDCTraitsFromClaims: map[string][]string{ - "email": {"superuser@example.com"}, - "groups": {"everyone", "idp-admin", "idp-dev"}, - "sub": {"00001234abcd"}, - }, - OIDCConnectorTraitMapping: []types.TraitMapping{ - { - Trait: "groups", - Value: "idp-admin", - Roles: []string{"access"}, - }, - }, - }, diagCtx.info) + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + ctx := context.Background() + s := setUpSuite(t) + // Create configurable IdP to use in tests. + idp := newFakeIDP(t, false /* tls */) + + // create role referenced in request. + role, err := types.NewRole("access", types.RoleSpecV5{ + Allow: types.RoleConditions{ + Logins: []string{"dummy"}, + }, + }) + require.NoError(t, err) + err = s.a.CreateRole(role) + require.NoError(t, err) + + // connector spec + spec := types.OIDCConnectorSpecV3{ + IssuerURL: idp.s.URL, + ClientID: "00000000000000000000000000000000", + ClientSecret: "0000000000000000000000000000000000000000000000000000000000000000", + Display: "Test", + Scope: []string{"groups"}, + ClaimsToRoles: tc.claimsToRoles, + RedirectURLs: []string{"https://proxy.example.com/v1/webapi/oidc/callback"}, + } + + oidcRequest := types.OIDCAuthRequest{ + ConnectorID: "-sso-test-okta", + Type: constants.OIDC, + CertTTL: defaults.OIDCAuthRequestTTL, + SSOTestFlow: true, + ConnectorSpec: &spec, + } + + request, err := s.a.CreateOIDCAuthRequest(ctx, oidcRequest) + require.NoError(t, err) + require.NotNil(t, request) + + values := url.Values{ + "code": []string{"XXX-code"}, + "state": []string{request.StateToken}, + } + + // override getClaimsFun. + s.a.getClaimsFun = func(closeCtx context.Context, oidcClient *oidc.Client, connector types.OIDCConnector, code string) (jose.Claims, error) { + cc := map[string]interface{}{ + "email_verified": true, + "groups": []string{"everyone", "idp-admin", "idp-dev"}, + "email": "superuser@example.com", + "sub": "00001234abcd", + "exp": 1652091713.0, + } + return cc, nil + } + + resp, err := s.a.ValidateOIDCAuthCallback(ctx, values) + if tc.wantValidateErr != nil { + require.ErrorIs(t, err, tc.wantValidateErr) + return + } + + require.NoError(t, err) + require.NotNil(t, resp) + require.Equal(t, &OIDCAuthResponse{ + Username: "superuser@example.com", + Identity: types.ExternalIdentity{ + ConnectorID: "-sso-test-okta", + Username: "superuser@example.com", + }, + Req: *request, + }, resp) + + diagCtx := ssoDiagContext{} + + resp, err = s.a.validateOIDCAuthCallback(ctx, &diagCtx, values) + require.NoError(t, err) + require.NotNil(t, resp) + require.Equal(t, &OIDCAuthResponse{ + Username: "superuser@example.com", + Identity: types.ExternalIdentity{ + ConnectorID: "-sso-test-okta", + Username: "superuser@example.com", + }, + Req: *request, + }, resp) + require.Equal(t, types.SSODiagnosticInfo{ + TestFlow: true, + Success: true, + CreateUserParams: &types.CreateUserParams{ + ConnectorName: "-sso-test-okta", + Username: "superuser@example.com", + Logins: nil, + KubeGroups: nil, + KubeUsers: nil, + Roles: []string{"access"}, + Traits: map[string][]string{ + "email": {"superuser@example.com"}, + "groups": {"everyone", "idp-admin", "idp-dev"}, + "sub": {"00001234abcd"}, + }, + SessionTTL: 600000000000, + }, + OIDCClaimsToRoles: []types.ClaimMapping{ + { + Claim: "groups", + Value: "idp-admin", + Roles: []string{"access"}, + }, + }, + OIDCClaimsToRolesWarnings: nil, + OIDCClaims: map[string]interface{}{ + "email_verified": true, + "groups": []string{"everyone", "idp-admin", "idp-dev"}, + "email": "superuser@example.com", + "sub": "00001234abcd", + "exp": 1652091713.0, + }, + OIDCIdentity: &types.OIDCIdentity{ + ID: "00001234abcd", + Name: "", + Email: "superuser@example.com", + ExpiresAt: diagCtx.info.OIDCIdentity.ExpiresAt, + }, + OIDCTraitsFromClaims: map[string][]string{ + "email": {"superuser@example.com"}, + "groups": {"everyone", "idp-admin", "idp-dev"}, + "sub": {"00001234abcd"}, + }, + OIDCConnectorTraitMapping: []types.TraitMapping{ + { + Trait: "groups", + Value: "idp-admin", + Roles: []string{"access"}, + }, + }, + }, diagCtx.info) + }) + } } // TestPingProvider confirms that the client_secret_post auth diff --git a/lib/auth/saml.go b/lib/auth/saml.go index f369e9af118..69e2fb10a47 100644 --- a/lib/auth/saml.go +++ b/lib/auth/saml.go @@ -40,6 +40,9 @@ import ( saml2 "github.com/russellhaering/gosaml2" ) +// ErrSAMLNoRoles results from not mapping any roles from SAML claims. +var ErrSAMLNoRoles = trace.AccessDenied("No roles mapped from claims. The mappings may contain typos.") + // UpsertSAMLConnector creates or updates a SAML connector. func (a *Server) UpsertSAMLConnector(ctx context.Context, connector types.SAMLConnector) error { if err := a.Identity.UpsertSAMLConnector(ctx, connector); err != nil { @@ -212,7 +215,7 @@ func (a *Server) calculateSAMLUser(diagCtx *ssoDiagContext, connector types.SAML Message: "No roles mapped for the user. The mappings may contain typos.", } } - return nil, trace.AccessDenied("No roles mapped from claims. The mappings may contain typos.") + return nil, trace.Wrap(ErrSAMLNoRoles) } // Pick smaller for role: session TTL from role or requested TTL. diff --git a/lib/client/redirect.go b/lib/client/redirect.go index bb68150ac16..fbf03ae1896 100644 --- a/lib/client/redirect.go +++ b/lib/client/redirect.go @@ -42,6 +42,10 @@ const ( // LoginFailedBadCallbackRedirectURL is a redirect URL when an SSO error specific to // auth connector's callback was encountered. LoginFailedBadCallbackRedirectURL = "/web/msg/error/login/callback" + + // LoginFailedUnauthorizedRedirectURL is a redirect URL for when an SSO authenticates successfully, + // but the user has no matching roles in Teleport. + LoginFailedUnauthorizedRedirectURL = "/web/msg/error/login/auth" ) // Redirector handles SSH redirect flow with the Teleport server diff --git a/lib/web/apiserver.go b/lib/web/apiserver.go index 9da37f2a7dd..8fdd450eb40 100644 --- a/lib/web/apiserver.go +++ b/lib/web/apiserver.go @@ -23,6 +23,7 @@ import ( "context" "encoding/base64" "encoding/json" + "errors" "fmt" "html/template" "io" @@ -1206,6 +1207,9 @@ func (h *Handler) githubCallback(w http.ResponseWriter, r *http.Request, p httpr } } } + if errors.Is(err, auth.ErrGithubNoTeams) { + return client.LoginFailedUnauthorizedRedirectURL + } return client.LoginFailedBadCallbackRedirectURL } @@ -1309,6 +1313,10 @@ func (h *Handler) oidcCallback(w http.ResponseWriter, r *http.Request, p httprou } } + if errors.Is(err, auth.ErrOIDCNoRoles) { + return client.LoginFailedUnauthorizedRedirectURL + } + return client.LoginFailedBadCallbackRedirectURL } diff --git a/lib/web/apiserver_test.go b/lib/web/apiserver_test.go index 9e3dd1f3689..12c3a27c129 100644 --- a/lib/web/apiserver_test.go +++ b/lib/web/apiserver_test.go @@ -543,110 +543,134 @@ func Test_clientMetaFromReq(t *testing.T) { }, got) } -func TestSAMLSuccess(t *testing.T) { +func TestSAML(t *testing.T) { t.Parallel() - ctx := context.Background() - s := newWebSuite(t) - input := fixtures.SAMLOktaConnectorV2 - decoder := kyaml.NewYAMLOrJSONDecoder(strings.NewReader(input), defaults.LookaheadBufSize) - var raw services.UnknownResource - err := decoder.Decode(&raw) - require.NoError(t, err) - - connector, err := services.UnmarshalSAMLConnector(raw.Raw) - require.NoError(t, err) - err = services.ValidateSAMLConnector(connector) - require.NoError(t, err) - - role, err := types.NewRoleV3(connector.GetAttributesToRoles()[0].Roles[0], types.RoleSpecV5{ - Options: types.RoleOptions{ - MaxSessionTTL: types.NewDuration(apidefaults.MaxCertDuration), + tests := []struct { + name string + rawConnector string + validSession bool + expectedRedirectURL string + }{ + { + name: "success", + rawConnector: fixtures.SAMLOktaConnectorV2, + validSession: true, + expectedRedirectURL: "/after", }, - Allow: types.RoleConditions{ - NodeLabels: types.Labels{types.Wildcard: []string{types.Wildcard}}, - Namespaces: []string{apidefaults.Namespace}, - Rules: []types.Rule{ - types.NewRule(types.Wildcard, services.RW()), - }, + { + name: "fail to map claims to roles", + rawConnector: strings.ReplaceAll(fixtures.SAMLOktaConnectorV2, "Everyone", "No-one"), + validSession: false, + expectedRedirectURL: client.LoginFailedUnauthorizedRedirectURL, }, - }) - require.NoError(t, err) - role.SetLogins(types.Allow, []string{s.user}) - err = s.server.Auth().UpsertRole(s.ctx, role) - require.NoError(t, err) + } - err = s.server.Auth().UpsertSAMLConnector(ctx, connector) - require.NoError(t, err) - s.server.Auth().SetClock(clockwork.NewFakeClockAt(time.Date(2017, 5, 10, 18, 53, 0, 0, time.UTC))) - clt := s.clientNoRedirects() + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + ctx := context.Background() + s := newWebSuite(t) + input := tc.rawConnector - csrfToken := "2ebcb768d0090ea4368e42880c970b61865c326172a4a2343b645cf5d7f20992" + decoder := kyaml.NewYAMLOrJSONDecoder(strings.NewReader(input), defaults.LookaheadBufSize) + var raw services.UnknownResource + err := decoder.Decode(&raw) + require.NoError(t, err) - baseURL, err := url.Parse(clt.Endpoint("webapi", "saml", "sso") + `?connector_id=` + connector.GetName() + `&redirect_url=http://localhost/after`) - require.NoError(t, err) - req, err := http.NewRequest("GET", baseURL.String(), nil) - require.NoError(t, err) - addCSRFCookieToReq(req, csrfToken) - re, err := clt.Client.RoundTrip(func() (*http.Response, error) { - return clt.Client.HTTPClient().Do(req) - }) - require.NoError(t, err) + connector, err := services.UnmarshalSAMLConnector(raw.Raw) + require.NoError(t, err) - // we got a redirect - urlPattern := regexp.MustCompile(`URL='([^']*)'`) - locationURL := urlPattern.FindStringSubmatch(string(re.Bytes()))[1] - u, err := url.Parse(locationURL) - require.NoError(t, err) - require.Equal(t, fixtures.SAMLOktaSSO, u.Scheme+"://"+u.Host+u.Path) - data, err := base64.StdEncoding.DecodeString(u.Query().Get("SAMLRequest")) - require.NoError(t, err) - buf, err := io.ReadAll(flate.NewReader(bytes.NewReader(data))) - require.NoError(t, err) - doc := etree.NewDocument() - err = doc.ReadFromBytes(buf) - require.NoError(t, err) - id := doc.Root().SelectAttr("ID") - require.NotNil(t, id) + role, err := types.NewRoleV3(connector.GetAttributesToRoles()[0].Roles[0], types.RoleSpecV5{ + Options: types.RoleOptions{ + MaxSessionTTL: types.NewDuration(apidefaults.MaxCertDuration), + }, + Allow: types.RoleConditions{ + NodeLabels: types.Labels{types.Wildcard: []string{types.Wildcard}}, + Namespaces: []string{apidefaults.Namespace}, + Rules: []types.Rule{ + types.NewRule(types.Wildcard, services.RW()), + }, + }, + }) + require.NoError(t, err) + role.SetLogins(types.Allow, []string{s.user}) + err = s.server.Auth().UpsertRole(s.ctx, role) + require.NoError(t, err) - authRequest, err := s.server.Auth().GetSAMLAuthRequest(context.Background(), id.Value) - require.NoError(t, err) + err = s.server.Auth().UpsertSAMLConnector(ctx, connector) + require.NoError(t, err) + s.server.Auth().SetClock(clockwork.NewFakeClockAt(time.Date(2017, 5, 10, 18, 53, 0, 0, time.UTC))) + clt := s.clientNoRedirects() - // now swap the request id to the hardcoded one in fixtures - authRequest.ID = fixtures.SAMLOktaAuthRequestID - authRequest.CSRFToken = csrfToken - err = s.server.Auth().Identity.CreateSAMLAuthRequest(ctx, *authRequest, backend.Forever) - require.NoError(t, err) + csrfToken := "2ebcb768d0090ea4368e42880c970b61865c326172a4a2343b645cf5d7f20992" - // now respond with pre-recorded request to the POST url - in := &bytes.Buffer{} - fw, err := flate.NewWriter(in, flate.DefaultCompression) - require.NoError(t, err) + baseURL, err := url.Parse(clt.Endpoint("webapi", "saml", "sso") + `?connector_id=` + connector.GetName() + `&redirect_url=http://localhost/after`) + require.NoError(t, err) + req, err := http.NewRequest("GET", baseURL.String(), nil) + require.NoError(t, err) + addCSRFCookieToReq(req, csrfToken) + re, err := clt.Client.RoundTrip(func() (*http.Response, error) { + return clt.Client.HTTPClient().Do(req) + }) + require.NoError(t, err) - _, err = fw.Write([]byte(fixtures.SAMLOktaAuthnResponseXML)) - require.NoError(t, err) - err = fw.Close() - require.NoError(t, err) - encodedResponse := base64.StdEncoding.EncodeToString(in.Bytes()) - require.NotNil(t, encodedResponse) + // we got a redirect + urlPattern := regexp.MustCompile(`URL='([^']*)'`) + locationURL := urlPattern.FindStringSubmatch(string(re.Bytes()))[1] + u, err := url.Parse(locationURL) + require.NoError(t, err) + require.Equal(t, fixtures.SAMLOktaSSO, u.Scheme+"://"+u.Host+u.Path) + data, err := base64.StdEncoding.DecodeString(u.Query().Get("SAMLRequest")) + require.NoError(t, err) + buf, err := io.ReadAll(flate.NewReader(bytes.NewReader(data))) + require.NoError(t, err) + doc := etree.NewDocument() + err = doc.ReadFromBytes(buf) + require.NoError(t, err) + id := doc.Root().SelectAttr("ID") + require.NotNil(t, id) - // now send the response to the server to exchange it for auth session - form := url.Values{} - form.Add("SAMLResponse", encodedResponse) - req, err = http.NewRequest("POST", clt.Endpoint("webapi", "saml", "acs"), strings.NewReader(form.Encode())) - req.Header.Add("Content-Type", "application/x-www-form-urlencoded") - addCSRFCookieToReq(req, csrfToken) - require.NoError(t, err) - authRe, err := clt.Client.RoundTrip(func() (*http.Response, error) { - return clt.Client.HTTPClient().Do(req) - }) + authRequest, err := s.server.Auth().GetSAMLAuthRequest(context.Background(), id.Value) + require.NoError(t, err) - require.NoError(t, err) - require.Equal(t, http.StatusFound, authRe.Code(), "Response: %v", string(authRe.Bytes())) - // we have got valid session - require.NotEmpty(t, authRe.Headers().Get("Set-Cookie")) - // we are being redirected to original URL - require.Equal(t, "/after", authRe.Headers().Get("Location")) + // now swap the request id to the hardcoded one in fixtures + authRequest.ID = fixtures.SAMLOktaAuthRequestID + authRequest.CSRFToken = csrfToken + err = s.server.Auth().Identity.CreateSAMLAuthRequest(ctx, *authRequest, backend.Forever) + require.NoError(t, err) + + // now respond with pre-recorded request to the POST url + in := &bytes.Buffer{} + fw, err := flate.NewWriter(in, flate.DefaultCompression) + require.NoError(t, err) + + _, err = fw.Write([]byte(fixtures.SAMLOktaAuthnResponseXML)) + require.NoError(t, err) + err = fw.Close() + require.NoError(t, err) + encodedResponse := base64.StdEncoding.EncodeToString(in.Bytes()) + require.NotNil(t, encodedResponse) + + // now send the response to the server to exchange it for auth session + form := url.Values{} + form.Add("SAMLResponse", encodedResponse) + req, err = http.NewRequest("POST", clt.Endpoint("webapi", "saml", "acs"), strings.NewReader(form.Encode())) + req.Header.Add("Content-Type", "application/x-www-form-urlencoded") + addCSRFCookieToReq(req, csrfToken) + require.NoError(t, err) + authRe, err := clt.Client.RoundTrip(func() (*http.Response, error) { + return clt.Client.HTTPClient().Do(req) + }) + + require.NoError(t, err) + require.Equal(t, http.StatusFound, authRe.Code(), "Response: %v", string(authRe.Bytes())) + if tc.validSession { + // we have got valid session + require.NotEmpty(t, authRe.Headers().Get("Set-Cookie")) + } + require.Equal(t, tc.expectedRedirectURL, authRe.Headers().Get("Location")) + }) + } } func TestWebSessionsCRUD(t *testing.T) { diff --git a/lib/web/saml.go b/lib/web/saml.go index a1120db01c5..7f0b12e4378 100644 --- a/lib/web/saml.go +++ b/lib/web/saml.go @@ -17,6 +17,7 @@ limitations under the License. package web import ( + "errors" "net/http" "github.com/gravitational/teleport/api/types" @@ -113,6 +114,10 @@ func (h *Handler) samlACS(w http.ResponseWriter, r *http.Request, p httprouter.P } } + if errors.Is(err, auth.ErrSAMLNoRoles) { + return client.LoginFailedUnauthorizedRedirectURL + } + return client.LoginFailedBadCallbackRedirectURL }