mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
Remove SAML connector and tests (#18285)
Remove the SAML connector from the OSS repository. It has been migrated to the enterprise repository.
This commit is contained in:
+1
-1
Submodule e updated: b37a41b87e...254b78ea7b
@@ -26,9 +26,6 @@ build_teleport_fuzzers() {
|
||||
compile_native_go_fuzzer $TELEPORT_PREFIX/lib/services \
|
||||
FuzzParserEvalBoolPredicate fuzz_parser_eval_bool_predicate
|
||||
|
||||
compile_native_go_fuzzer $TELEPORT_PREFIX/lib/auth \
|
||||
FuzzParseSAMLInResponseTo fuzz_parse_saml_in_response_to
|
||||
|
||||
compile_native_go_fuzzer $TELEPORT_PREFIX/lib/restrictedsession \
|
||||
FuzzParseIPSpec fuzz_parse_ip_spec
|
||||
|
||||
|
||||
@@ -290,13 +290,6 @@ func NewServer(cfg *InitConfig, opts ...ServerOption) (*Server, error) {
|
||||
)
|
||||
}
|
||||
}
|
||||
// Plug in SAML auth service
|
||||
sas := NewSAMLAuthService(&SAMLAuthServiceConfig{
|
||||
Auth: &as,
|
||||
Emitter: as.emitter,
|
||||
AssertionReplayService: as.Unstable.AssertionReplayService,
|
||||
})
|
||||
as.SetSAMLService(sas)
|
||||
|
||||
return &as, nil
|
||||
}
|
||||
|
||||
@@ -44,7 +44,6 @@ import (
|
||||
"github.com/gravitational/teleport/lib/auth/testauthority"
|
||||
libdefaults "github.com/gravitational/teleport/lib/defaults"
|
||||
"github.com/gravitational/teleport/lib/events"
|
||||
"github.com/gravitational/teleport/lib/fixtures"
|
||||
"github.com/gravitational/teleport/lib/services"
|
||||
"github.com/gravitational/teleport/lib/session"
|
||||
"github.com/gravitational/teleport/lib/tlsca"
|
||||
@@ -156,181 +155,6 @@ func TestSSOUserCanReissueCert(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestSAMLAuthRequest(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
srv := newTestTLSServer(t)
|
||||
|
||||
emptyRole, err := CreateRole(ctx, srv.Auth(), "test-empty", types.RoleSpecV5{})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = CreateRole(ctx, srv.Auth(), "baz", types.RoleSpecV5{})
|
||||
require.NoError(t, err)
|
||||
|
||||
access1Role, err := CreateRole(ctx, srv.Auth(), "test-access-1", types.RoleSpecV5{
|
||||
Allow: types.RoleConditions{
|
||||
Rules: []types.Rule{
|
||||
{
|
||||
Resources: []string{types.KindSAMLRequest},
|
||||
Verbs: []string{types.VerbCreate},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
access2Role, err := CreateRole(ctx, srv.Auth(), "test-access-2", types.RoleSpecV5{
|
||||
Allow: types.RoleConditions{
|
||||
Rules: []types.Rule{
|
||||
{
|
||||
Resources: []string{types.KindSAML},
|
||||
Verbs: []string{types.VerbCreate},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
access3Role, err := CreateRole(ctx, srv.Auth(), "test-access-3", types.RoleSpecV5{
|
||||
Allow: types.RoleConditions{
|
||||
Rules: []types.Rule{
|
||||
{
|
||||
Resources: []string{types.KindSAML, types.KindSAMLRequest},
|
||||
Verbs: []string{types.VerbCreate},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
readerRole, err := CreateRole(ctx, srv.Auth(), "test-access-4", types.RoleSpecV5{
|
||||
Allow: types.RoleConditions{
|
||||
Rules: []types.Rule{
|
||||
{
|
||||
Resources: []string{types.KindSAMLRequest},
|
||||
Verbs: []string{types.VerbRead},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
conn, err := types.NewSAMLConnector("foo", types.SAMLConnectorSpecV2{
|
||||
Issuer: "test",
|
||||
SSO: "test",
|
||||
Cert: fixtures.TLSCACertPEM,
|
||||
AssertionConsumerService: "test",
|
||||
AttributesToRoles: []types.AttributeMapping{{
|
||||
Name: "foo",
|
||||
Value: "bar",
|
||||
Roles: []string{"baz"},
|
||||
}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
err = srv.Auth().UpsertSAMLConnector(ctx, conn)
|
||||
require.NoError(t, err)
|
||||
|
||||
reqNormal := types.SAMLAuthRequest{ConnectorID: conn.GetName(), Type: constants.SAML}
|
||||
reqTest := types.SAMLAuthRequest{ConnectorID: conn.GetName(), Type: constants.SAML, SSOTestFlow: true, ConnectorSpec: &types.SAMLConnectorSpecV2{
|
||||
Issuer: "test",
|
||||
Audience: "test",
|
||||
ServiceProviderIssuer: "test",
|
||||
SSO: "test",
|
||||
Cert: fixtures.TLSCACertPEM,
|
||||
AssertionConsumerService: "test",
|
||||
AttributesToRoles: []types.AttributeMapping{{
|
||||
Name: "foo",
|
||||
Value: "bar",
|
||||
Roles: []string{"baz"},
|
||||
}},
|
||||
}}
|
||||
|
||||
tests := []struct {
|
||||
desc string
|
||||
roles []string
|
||||
request types.SAMLAuthRequest
|
||||
expectAccessDenied bool
|
||||
}{
|
||||
{
|
||||
desc: "empty role - no access",
|
||||
roles: []string{emptyRole.GetName()},
|
||||
request: reqNormal,
|
||||
expectAccessDenied: true,
|
||||
},
|
||||
{
|
||||
desc: "can create regular request with normal access",
|
||||
roles: []string{access1Role.GetName()},
|
||||
request: reqNormal,
|
||||
expectAccessDenied: false,
|
||||
},
|
||||
{
|
||||
desc: "cannot create sso test request with normal access",
|
||||
roles: []string{access1Role.GetName()},
|
||||
request: reqTest,
|
||||
expectAccessDenied: true,
|
||||
},
|
||||
{
|
||||
desc: "cannot create normal request with connector access",
|
||||
roles: []string{access2Role.GetName()},
|
||||
request: reqNormal,
|
||||
expectAccessDenied: true,
|
||||
},
|
||||
{
|
||||
desc: "cannot create sso test request with connector access",
|
||||
roles: []string{access2Role.GetName()},
|
||||
request: reqTest,
|
||||
expectAccessDenied: true,
|
||||
},
|
||||
{
|
||||
desc: "can create regular request with combined access",
|
||||
roles: []string{access3Role.GetName()},
|
||||
request: reqNormal,
|
||||
expectAccessDenied: false,
|
||||
},
|
||||
{
|
||||
desc: "can create sso test request with combined access",
|
||||
roles: []string{access3Role.GetName()},
|
||||
request: reqTest,
|
||||
expectAccessDenied: false,
|
||||
},
|
||||
}
|
||||
|
||||
user, err := CreateUser(srv.Auth(), "dummy")
|
||||
require.NoError(t, err)
|
||||
|
||||
userReader, err := CreateUser(srv.Auth(), "dummy-reader", readerRole)
|
||||
require.NoError(t, err)
|
||||
|
||||
clientReader, err := srv.NewClient(TestUser(userReader.GetName()))
|
||||
require.NoError(t, err)
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.desc, func(t *testing.T) {
|
||||
user.SetRoles(tt.roles)
|
||||
err = srv.Auth().UpsertUser(user)
|
||||
require.NoError(t, err)
|
||||
|
||||
client, err := srv.NewClient(TestUser(user.GetName()))
|
||||
require.NoError(t, err)
|
||||
|
||||
request, err := client.CreateSAMLAuthRequest(ctx, tt.request)
|
||||
if tt.expectAccessDenied {
|
||||
require.Error(t, err)
|
||||
require.True(t, trace.IsAccessDenied(err), "expected access denied, got: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, request.ID)
|
||||
require.Equal(t, tt.request.ConnectorID, request.ConnectorID)
|
||||
|
||||
requestCopy, err := clientReader.GetSAMLAuthRequest(ctx, request.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, request, requestCopy)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestInstaller(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
srv := newTestTLSServer(t)
|
||||
|
||||
@@ -17,24 +17,11 @@ limitations under the License.
|
||||
package auth
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"testing"
|
||||
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func FuzzParseSAMLInResponseTo(f *testing.F) {
|
||||
// Disable Go App Engine logging
|
||||
logrus.SetLevel(logrus.PanicLevel)
|
||||
|
||||
f.Fuzz(func(t *testing.T, response string) {
|
||||
require.NotPanics(t, func() {
|
||||
ParseSAMLInResponseTo(base64.StdEncoding.EncodeToString([]byte(response)))
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
func FuzzParseAndVerifyIID(f *testing.F) {
|
||||
f.Fuzz(func(t *testing.T, iidBytes []byte) {
|
||||
require.NotPanics(t, func() {
|
||||
|
||||
@@ -17,31 +17,16 @@ limitations under the License.
|
||||
package auth
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"compress/flate"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"sync"
|
||||
|
||||
"github.com/beevik/etree"
|
||||
"github.com/google/go-cmp/cmp"
|
||||
"github.com/gravitational/trace"
|
||||
saml2 "github.com/russellhaering/gosaml2"
|
||||
|
||||
"github.com/gravitational/teleport"
|
||||
"github.com/gravitational/teleport/api/constants"
|
||||
apidefaults "github.com/gravitational/teleport/api/defaults"
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
apievents "github.com/gravitational/teleport/api/types/events"
|
||||
"github.com/gravitational/teleport/api/utils/keys"
|
||||
"github.com/gravitational/teleport/lib/defaults"
|
||||
"github.com/gravitational/teleport/lib/events"
|
||||
"github.com/gravitational/teleport/lib/services"
|
||||
"github.com/gravitational/teleport/lib/services/local"
|
||||
"github.com/gravitational/teleport/lib/utils"
|
||||
)
|
||||
|
||||
// ErrSAMLRequiresEnterprise is the error returned by the SAML methods when not
|
||||
@@ -129,307 +114,6 @@ func (a *Server) ValidateSAMLResponse(ctx context.Context, re string, connectorI
|
||||
return resp, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// SAMLAuthService implements the logic of the SAML connector, allowing SSO
|
||||
// logins using the SAML protocol.
|
||||
//
|
||||
// SAMLAuthService implements the SAMLService interface.
|
||||
type SAMLAuthService struct {
|
||||
auth *Server
|
||||
emitter apievents.Emitter
|
||||
assertionReplayService *local.AssertionReplayService
|
||||
samlProviders map[string]*samlProvider
|
||||
lock sync.Mutex
|
||||
}
|
||||
|
||||
type SAMLAuthServiceConfig struct {
|
||||
Auth *Server
|
||||
Emitter apievents.Emitter
|
||||
AssertionReplayService *local.AssertionReplayService
|
||||
}
|
||||
|
||||
// NewSAMLAuthService returns a SAMLAuthService configured to use the
|
||||
// services given in the config.
|
||||
func NewSAMLAuthService(cfg *SAMLAuthServiceConfig) *SAMLAuthService {
|
||||
return &SAMLAuthService{
|
||||
auth: cfg.Auth,
|
||||
emitter: cfg.Emitter,
|
||||
assertionReplayService: cfg.AssertionReplayService,
|
||||
|
||||
samlProviders: make(map[string]*samlProvider),
|
||||
}
|
||||
}
|
||||
|
||||
// samlProvider is internal structure that stores SAML client and its config
|
||||
type samlProvider struct {
|
||||
provider *saml2.SAMLServiceProvider
|
||||
connector types.SAMLConnector
|
||||
}
|
||||
|
||||
// ErrSAMLNoRoles results from not mapping any roles from SAML claims.
|
||||
var ErrSAMLNoRoles = trace.AccessDenied("No roles mapped from claims. The mappings may contain typos.")
|
||||
|
||||
func (sas *SAMLAuthService) CreateSAMLAuthRequest(ctx context.Context, req types.SAMLAuthRequest) (*types.SAMLAuthRequest, error) {
|
||||
connector, provider, err := sas.getSAMLConnectorAndProvider(ctx, req)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
doc, err := provider.BuildAuthRequestDocument()
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
attr := doc.Root().SelectAttr("ID")
|
||||
if attr == nil || attr.Value == "" {
|
||||
return nil, trace.BadParameter("missing auth request ID")
|
||||
}
|
||||
|
||||
req.ID = attr.Value
|
||||
|
||||
// Workaround for Ping: Ping expects `SigAlg` and `Signature` query
|
||||
// parameters when "Enforce Signed Authn Request" is enabled, but gosaml2
|
||||
// only provides these parameters when binding == BindingHttpRedirect.
|
||||
// Luckily, BuildAuthURLRedirect sets this and is otherwise identical to
|
||||
// the standard BuildAuthURLFromDocument.
|
||||
if connector.GetProvider() == teleport.Ping {
|
||||
req.RedirectURL, err = provider.BuildAuthURLRedirect("", doc)
|
||||
} else {
|
||||
req.RedirectURL, err = provider.BuildAuthURLFromDocument("", doc)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
err = sas.auth.Services.CreateSAMLAuthRequest(ctx, req, defaults.SAMLAuthRequestTTL)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
return &req, nil
|
||||
}
|
||||
|
||||
func (sas *SAMLAuthService) getSAMLConnectorAndProviderByID(ctx context.Context, connectorID string) (types.SAMLConnector, *saml2.SAMLServiceProvider, error) {
|
||||
connector, err := sas.auth.Identity.GetSAMLConnector(ctx, connectorID, true)
|
||||
if err != nil {
|
||||
return nil, nil, trace.Wrap(err)
|
||||
}
|
||||
provider, err := sas.getSAMLProvider(connector)
|
||||
if err != nil {
|
||||
return nil, nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
return connector, provider, nil
|
||||
}
|
||||
|
||||
func (sas *SAMLAuthService) getSAMLConnectorAndProvider(ctx context.Context, req types.SAMLAuthRequest) (types.SAMLConnector, *saml2.SAMLServiceProvider, error) {
|
||||
if req.SSOTestFlow {
|
||||
if req.ConnectorSpec == nil {
|
||||
return nil, nil, trace.BadParameter("ConnectorSpec cannot be nil when SSOTestFlow is true")
|
||||
}
|
||||
|
||||
if req.ConnectorID == "" {
|
||||
return nil, nil, trace.BadParameter("ConnectorID cannot be empty")
|
||||
}
|
||||
|
||||
// stateless test flow
|
||||
connector, err := types.NewSAMLConnector(req.ConnectorID, *req.ConnectorSpec)
|
||||
if err != nil {
|
||||
return nil, nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// validate, set defaults for connector
|
||||
err = services.ValidateSAMLConnector(connector, sas.auth)
|
||||
if err != nil {
|
||||
return nil, nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// we don't want to cache the provider. construct it directly instead of using sas.getSAMLProvider()
|
||||
provider, err := services.GetSAMLServiceProvider(connector, sas.auth.GetClock())
|
||||
if err != nil {
|
||||
return nil, nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
return connector, provider, nil
|
||||
}
|
||||
|
||||
// regular execution flow
|
||||
return sas.getSAMLConnectorAndProviderByID(ctx, req.ConnectorID)
|
||||
}
|
||||
|
||||
func (sas *SAMLAuthService) getSAMLProvider(conn types.SAMLConnector) (*saml2.SAMLServiceProvider, error) {
|
||||
sas.lock.Lock()
|
||||
defer sas.lock.Unlock()
|
||||
|
||||
providerPack, ok := sas.samlProviders[conn.GetName()]
|
||||
if ok && cmp.Equal(providerPack.connector, conn) {
|
||||
return providerPack.provider, nil
|
||||
}
|
||||
delete(sas.samlProviders, conn.GetName())
|
||||
|
||||
serviceProvider, err := services.GetSAMLServiceProvider(conn, sas.auth.GetClock())
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
sas.samlProviders[conn.GetName()] = &samlProvider{connector: conn, provider: serviceProvider}
|
||||
|
||||
return serviceProvider, nil
|
||||
}
|
||||
|
||||
func (sas *SAMLAuthService) calculateSAMLUser(diagCtx *SSODiagContext, connector types.SAMLConnector, assertionInfo saml2.AssertionInfo, request *types.SAMLAuthRequest) (*CreateUserParams, error) {
|
||||
p := CreateUserParams{
|
||||
ConnectorName: connector.GetName(),
|
||||
Username: assertionInfo.NameID,
|
||||
}
|
||||
|
||||
p.Traits = services.SAMLAssertionsToTraits(assertionInfo)
|
||||
|
||||
diagCtx.Info.SAMLTraitsFromAssertions = p.Traits
|
||||
diagCtx.Info.SAMLConnectorTraitMapping = connector.GetTraitMappings()
|
||||
|
||||
var warnings []string
|
||||
warnings, p.Roles = services.TraitsToRoles(connector.GetTraitMappings(), p.Traits)
|
||||
if len(p.Roles) == 0 {
|
||||
if len(warnings) != 0 {
|
||||
log.WithField("connector", connector).Warnf("No roles mapped from claims. Warnings: %q", warnings)
|
||||
diagCtx.Info.SAMLAttributesToRolesWarnings = &types.SSOWarnings{
|
||||
Message: "No roles mapped for the user",
|
||||
Warnings: warnings,
|
||||
}
|
||||
} else {
|
||||
log.WithField("connector", connector).Warnf("No roles mapped from claims.")
|
||||
diagCtx.Info.SAMLAttributesToRolesWarnings = &types.SSOWarnings{
|
||||
Message: "No roles mapped for the user. The mappings may contain typos.",
|
||||
}
|
||||
}
|
||||
return nil, trace.Wrap(ErrSAMLNoRoles)
|
||||
}
|
||||
|
||||
// Pick smaller for role: session TTL from role or requested TTL.
|
||||
roles, err := services.FetchRoles(p.Roles, sas.auth, p.Traits)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
roleTTL := roles.AdjustSessionTTL(apidefaults.MaxCertDuration)
|
||||
|
||||
if request != nil {
|
||||
p.SessionTTL = utils.MinTTL(roleTTL, request.CertTTL)
|
||||
} else {
|
||||
p.SessionTTL = roleTTL
|
||||
}
|
||||
|
||||
return &p, nil
|
||||
}
|
||||
|
||||
func (sas *SAMLAuthService) createSAMLUser(p *CreateUserParams, dryRun bool) (types.User, error) {
|
||||
expires := sas.auth.GetClock().Now().UTC().Add(p.SessionTTL)
|
||||
|
||||
log.Debugf("Generating dynamic SAML identity %v/%v with roles: %v. Dry run: %v.", p.ConnectorName, p.Username, p.Roles, dryRun)
|
||||
|
||||
user := &types.UserV2{
|
||||
Kind: types.KindUser,
|
||||
Version: types.V2,
|
||||
Metadata: types.Metadata{
|
||||
Name: p.Username,
|
||||
Namespace: apidefaults.Namespace,
|
||||
Expires: &expires,
|
||||
},
|
||||
Spec: types.UserSpecV2{
|
||||
Roles: p.Roles,
|
||||
Traits: p.Traits,
|
||||
SAMLIdentities: []types.ExternalIdentity{
|
||||
{
|
||||
ConnectorID: p.ConnectorName,
|
||||
Username: p.Username,
|
||||
},
|
||||
},
|
||||
CreatedBy: types.CreatedBy{
|
||||
User: types.UserRef{
|
||||
Name: teleport.UserSystem,
|
||||
},
|
||||
Time: sas.auth.GetClock().Now().UTC(),
|
||||
Connector: &types.ConnectorRef{
|
||||
Type: constants.SAML,
|
||||
ID: p.ConnectorName,
|
||||
Identity: p.Username,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
if dryRun {
|
||||
return user, nil
|
||||
}
|
||||
|
||||
// Get the user to check if it already exists or not.
|
||||
existingUser, err := sas.auth.Services.GetUser(p.Username, false)
|
||||
if err != nil && !trace.IsNotFound(err) {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
ctx := context.TODO()
|
||||
|
||||
// Overwrite exisiting user if it was created from an external identity provider.
|
||||
if existingUser != nil {
|
||||
connectorRef := existingUser.GetCreatedBy().Connector
|
||||
|
||||
// If the exisiting user is a local user, fail and advise how to fix the problem.
|
||||
if connectorRef == nil {
|
||||
return nil, trace.AlreadyExists("local user with name %q already exists. Either change "+
|
||||
"NameID in assertion or remove local user and try again.", existingUser.GetName())
|
||||
}
|
||||
|
||||
log.Debugf("Overwriting existing user %q created with %v connector %v.",
|
||||
existingUser.GetName(), connectorRef.Type, connectorRef.ID)
|
||||
|
||||
if err := sas.auth.UpdateUser(ctx, user); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
} else {
|
||||
if err := sas.auth.CreateUser(ctx, user); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
}
|
||||
|
||||
return user, nil
|
||||
}
|
||||
|
||||
func ParseSAMLInResponseTo(response string) (string, error) {
|
||||
raw, _ := base64.StdEncoding.DecodeString(response)
|
||||
|
||||
doc := etree.NewDocument()
|
||||
err := doc.ReadFromBytes(raw)
|
||||
if err != nil {
|
||||
// Attempt to inflate the response in case it happens to be compressed (as with one case at saml.oktadev.com)
|
||||
buf, err := io.ReadAll(flate.NewReader(bytes.NewReader(raw)))
|
||||
if err != nil {
|
||||
return "", trace.Wrap(err)
|
||||
}
|
||||
|
||||
doc = etree.NewDocument()
|
||||
err = doc.ReadFromBytes(buf)
|
||||
if err != nil {
|
||||
return "", trace.Wrap(err)
|
||||
}
|
||||
}
|
||||
|
||||
if doc.Root() == nil {
|
||||
return "", trace.BadParameter("unable to parse response")
|
||||
}
|
||||
|
||||
// Try to find the InResponseTo attribute in the SAML response. If we can't find this, return
|
||||
// a predictable error message so the caller may choose interpret it as an IdP-initiated payload.
|
||||
el := doc.Root()
|
||||
responseTo := el.SelectAttr("InResponseTo")
|
||||
if responseTo == nil {
|
||||
return "", trace.NotFound("missing InResponseTo attribute")
|
||||
}
|
||||
if responseTo.Value == "" {
|
||||
return "", trace.BadParameter("InResponseTo can not be empty")
|
||||
}
|
||||
return responseTo.Value, nil
|
||||
}
|
||||
|
||||
// SAMLAuthResponse is returned when auth server validated callback parameters
|
||||
// returned from SAML identity provider
|
||||
type SAMLAuthResponse struct {
|
||||
@@ -494,256 +178,3 @@ type SAMLAuthRawResponse struct {
|
||||
// TLSCert is TLS certificate authority certificate
|
||||
TLSCert []byte `json:"tls_cert,omitempty"`
|
||||
}
|
||||
|
||||
// SAMLAuthRequestFromProto converts the types.SAMLAuthRequest to SAMLAuthRequestData.
|
||||
func SAMLAuthRequestFromProto(req *types.SAMLAuthRequest) SAMLAuthRequest {
|
||||
return SAMLAuthRequest{
|
||||
ID: req.ID,
|
||||
PublicKey: req.PublicKey,
|
||||
CSRFToken: req.CSRFToken,
|
||||
CreateWebSession: req.CreateWebSession,
|
||||
ClientRedirectURL: req.ClientRedirectURL,
|
||||
}
|
||||
}
|
||||
|
||||
// ValidateSAMLResponse consumes attribute statements from SAML identity provider
|
||||
func (sas *SAMLAuthService) ValidateSAMLResponse(ctx context.Context, samlResponse string, connectorID string) (*SAMLAuthResponse, error) {
|
||||
event := &apievents.UserLogin{
|
||||
Metadata: apievents.Metadata{
|
||||
Type: events.UserLoginEvent,
|
||||
},
|
||||
Method: events.LoginMethodSAML,
|
||||
}
|
||||
|
||||
diagCtx := NewSSODiagContext(types.KindSAML, sas.auth)
|
||||
|
||||
auth, err := sas.validateSAMLResponse(ctx, diagCtx, samlResponse, connectorID)
|
||||
diagCtx.Info.Error = trace.UserMessage(err)
|
||||
|
||||
diagCtx.WriteToBackend(ctx)
|
||||
|
||||
attributeStatements := diagCtx.Info.SAMLAttributeStatements
|
||||
if attributeStatements != nil {
|
||||
attributes, err := apievents.EncodeMapStrings(attributeStatements)
|
||||
if err != nil {
|
||||
event.Status.UserMessage = fmt.Sprintf("Failed to encode identity attributes: %v", err.Error())
|
||||
log.WithError(err).Debug("Failed to encode identity attributes.")
|
||||
} else {
|
||||
event.IdentityAttributes = attributes
|
||||
}
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
event.Code = events.UserSSOLoginFailureCode
|
||||
if diagCtx.Info.TestFlow {
|
||||
event.Code = events.UserSSOTestFlowLoginFailureCode
|
||||
}
|
||||
event.Status.Success = false
|
||||
event.Status.Error = trace.Unwrap(err).Error()
|
||||
event.Status.UserMessage = err.Error()
|
||||
if err := sas.emitter.EmitAuditEvent(ctx, event); err != nil {
|
||||
log.WithError(err).Warn("Failed to emit SAML login failed event.")
|
||||
}
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
event.Status.Success = true
|
||||
event.User = auth.Username
|
||||
event.Code = events.UserSSOLoginCode
|
||||
if diagCtx.Info.TestFlow {
|
||||
event.Code = events.UserSSOTestFlowLoginCode
|
||||
}
|
||||
|
||||
if err := sas.emitter.EmitAuditEvent(ctx, event); err != nil {
|
||||
log.WithError(err).Warn("Failed to emit SAML login event.")
|
||||
}
|
||||
|
||||
return auth, nil
|
||||
}
|
||||
|
||||
func (sas *SAMLAuthService) checkIDPInitiatedSAML(ctx context.Context, connector types.SAMLConnector, assertion *saml2.AssertionInfo) error {
|
||||
if !connector.GetAllowIDPInitiated() {
|
||||
return trace.AccessDenied("IdP initiated SAML is not allowed by the connector configuration")
|
||||
}
|
||||
|
||||
// Not all IdP's provide these variables, replay mitigation is best effort.
|
||||
if assertion.SessionIndex != "" || assertion.SessionNotOnOrAfter == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
err := sas.assertionReplayService.RecognizeSSOAssertion(ctx, connector.GetName(), assertion.SessionIndex, assertion.NameID, *assertion.SessionNotOnOrAfter)
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
|
||||
func (sas *SAMLAuthService) validateSAMLResponse(ctx context.Context, diagCtx *SSODiagContext, samlResponse string, connectorID string) (*SAMLAuthResponse, error) {
|
||||
idpInitiated := false
|
||||
var connector types.SAMLConnector
|
||||
var provider *saml2.SAMLServiceProvider
|
||||
var request *types.SAMLAuthRequest
|
||||
requestID, err := ParseSAMLInResponseTo(samlResponse)
|
||||
switch {
|
||||
case trace.IsNotFound(err):
|
||||
if connectorID == "" {
|
||||
return nil, trace.BadParameter("ACS URI did not include a valid SAML connector ID parameter")
|
||||
}
|
||||
|
||||
idpInitiated = true
|
||||
connector, provider, err = sas.getSAMLConnectorAndProviderByID(ctx, connectorID)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err, "Failed to get SAML connector and provider")
|
||||
}
|
||||
case err != nil:
|
||||
return nil, trace.Wrap(err)
|
||||
default:
|
||||
diagCtx.RequestID = requestID
|
||||
request, err = sas.auth.Identity.GetSAMLAuthRequest(ctx, requestID)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err, "Failed to get SAML Auth Request")
|
||||
}
|
||||
|
||||
diagCtx.Info.TestFlow = request.SSOTestFlow
|
||||
connector, provider, err = sas.getSAMLConnectorAndProvider(ctx, *request)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err, "Failed to get SAML connector and provider")
|
||||
}
|
||||
}
|
||||
|
||||
assertionInfo, err := provider.RetrieveAssertionInfo(samlResponse)
|
||||
if err != nil {
|
||||
return nil, trace.AccessDenied("received response with incorrect or missing attribute statements, please check the identity provider configuration to make sure that mappings for claims/attribute statements are set up correctly. <See: https://goteleport.com/teleport/docs/enterprise/sso/ssh-sso/>, failed to retrieve SAML assertion info from response: %v.", err).AddUserMessage("Failed to retrieve assertion info. This may indicate IdP configuration error.")
|
||||
}
|
||||
|
||||
if assertionInfo != nil {
|
||||
diagCtx.Info.SAMLAssertionInfo = (*types.AssertionInfo)(assertionInfo)
|
||||
}
|
||||
|
||||
if idpInitiated {
|
||||
if err := sas.checkIDPInitiatedSAML(ctx, connector, assertionInfo); err != nil {
|
||||
if trace.IsAccessDenied(err) {
|
||||
log.Warnf("Failed to process IdP-initiated login request. IdP-initiated login is disabled for this connector: %v.", err)
|
||||
}
|
||||
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
}
|
||||
|
||||
if assertionInfo.WarningInfo.InvalidTime {
|
||||
return nil, trace.AccessDenied("invalid time in SAML assertion info").AddUserMessage("SAML assertion info contained warning: invalid time.")
|
||||
}
|
||||
|
||||
if assertionInfo.WarningInfo.NotInAudience {
|
||||
return nil, trace.AccessDenied("no audience in SAML assertion info").AddUserMessage("SAML: not in expected audience. Check auth connector audience field and IdP configuration for typos and other errors.")
|
||||
}
|
||||
|
||||
log.Debugf("Obtained SAML assertions for %q.", assertionInfo.NameID)
|
||||
log.Debugf("SAML assertion warnings: %+v.", assertionInfo.WarningInfo)
|
||||
|
||||
attributeStatements := map[string][]string{}
|
||||
|
||||
for key, val := range assertionInfo.Values {
|
||||
var vals []string
|
||||
for _, vv := range val.Values {
|
||||
vals = append(vals, vv.Value)
|
||||
}
|
||||
log.Debugf("SAML assertion: %q: %q.", key, vals)
|
||||
attributeStatements[key] = vals
|
||||
}
|
||||
|
||||
diagCtx.Info.SAMLAttributeStatements = attributeStatements
|
||||
diagCtx.Info.SAMLAttributesToRoles = connector.GetAttributesToRoles()
|
||||
|
||||
if len(connector.GetAttributesToRoles()) == 0 {
|
||||
return nil, trace.BadParameter("no attributes to roles mapping, check connector documentation").AddUserMessage("Attributes-to-roles mapping is empty, SSO user will never have any roles.")
|
||||
}
|
||||
|
||||
log.Debugf("Applying %v SAML attribute to roles mappings.", len(connector.GetAttributesToRoles()))
|
||||
|
||||
// Calculate (figure out name, roles, traits, session TTL) of user and
|
||||
// create the user in the backend.
|
||||
params, err := sas.calculateSAMLUser(diagCtx, connector, *assertionInfo, request)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err, "Failed to calculate user attributes.")
|
||||
}
|
||||
|
||||
diagCtx.Info.CreateUserParams = &types.CreateUserParams{
|
||||
ConnectorName: params.ConnectorName,
|
||||
Username: params.Username,
|
||||
KubeGroups: params.KubeGroups,
|
||||
KubeUsers: params.KubeUsers,
|
||||
Roles: params.Roles,
|
||||
Traits: params.Traits,
|
||||
SessionTTL: types.Duration(params.SessionTTL),
|
||||
}
|
||||
|
||||
user, err := sas.createSAMLUser(params, request != nil && request.SSOTestFlow)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err, "Failed to create user from provided parameters.")
|
||||
}
|
||||
|
||||
// Auth was successful, return session, certificate, etc. to caller.
|
||||
auth := &SAMLAuthResponse{
|
||||
Identity: types.ExternalIdentity{
|
||||
ConnectorID: params.ConnectorName,
|
||||
Username: params.Username,
|
||||
},
|
||||
Username: user.GetName(),
|
||||
}
|
||||
|
||||
if request != nil {
|
||||
auth.Req = SAMLAuthRequestFromProto(request)
|
||||
} else {
|
||||
auth.Req = SAMLAuthRequest{
|
||||
CreateWebSession: true,
|
||||
}
|
||||
}
|
||||
|
||||
// In test flow skip signing and creating web sessions.
|
||||
if request != nil && request.SSOTestFlow {
|
||||
diagCtx.Info.Success = true
|
||||
return auth, nil
|
||||
}
|
||||
|
||||
// If the request is coming from a browser, create a web session.
|
||||
if request == nil || request.CreateWebSession {
|
||||
session, err := sas.auth.CreateWebSessionFromReq(ctx, types.NewWebSessionRequest{
|
||||
User: user.GetName(),
|
||||
Roles: user.GetRoles(),
|
||||
Traits: user.GetTraits(),
|
||||
SessionTTL: params.SessionTTL,
|
||||
LoginTime: sas.auth.GetClock().Now().UTC(),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err, "Failed to create web session.")
|
||||
}
|
||||
|
||||
auth.Session = session
|
||||
}
|
||||
|
||||
// If a public key was provided, sign it and return a certificate.
|
||||
if request != nil && len(request.PublicKey) != 0 {
|
||||
sshCert, tlsCert, err := sas.auth.CreateSessionCert(user, params.SessionTTL, request.PublicKey, request.Compatibility, request.RouteToCluster,
|
||||
request.KubernetesCluster, keys.AttestationStatementFromProto(request.AttestationStatement))
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err, "Failed to create session certificate.")
|
||||
}
|
||||
clusterName, err := sas.auth.GetClusterName()
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err, "Failed to obtain cluster name.")
|
||||
}
|
||||
auth.Cert = sshCert
|
||||
auth.TLSCert = tlsCert
|
||||
|
||||
// Return the host CA for this cluster only.
|
||||
authority, err := sas.auth.GetCertAuthority(ctx, types.CertAuthID{
|
||||
Type: types.HostCA,
|
||||
DomainName: clusterName.GetClusterName(),
|
||||
}, false)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err, "Failed to obtain cluster's host CA.")
|
||||
}
|
||||
auth.HostSigners = append(auth.HostSigners, authority)
|
||||
}
|
||||
|
||||
diagCtx.Info.Success = true
|
||||
return auth, nil
|
||||
}
|
||||
|
||||
File diff suppressed because one or more lines are too long
Reference in New Issue
Block a user