Add support for IdP-Initiated SAML2 login (#13924)

This commit is contained in:
Joel
2022-08-22 18:27:44 +00:00
committed by GitHub
parent 253173376f
commit 6caba42ec1
19 changed files with 1222 additions and 975 deletions
+2 -2
View File
@@ -7370,12 +7370,12 @@ func (m *UpgradeWindowStartUpdate) XXX_DiscardUnknown() {
var xxx_messageInfo_UpgradeWindowStartUpdate proto.InternalMessageInfo
// SessionRecordingAccess is emitted when a session is viewed in the web UI, allowing
// SessionRecordingAccess is emitted when a session recording is accessed, allowing
// session views to be included in the audit log
type SessionRecordingAccess struct {
// Metadata is a common event metadata.
Metadata `protobuf:"bytes,1,opt,name=Metadata,proto3,embedded=Metadata" json:""`
// SessionID is the ID of the application session.
// SessionID is the ID of the session.
SessionID string `protobuf:"bytes,2,opt,name=SessionID,proto3" json:"sid"`
// UserMetadata is a common user event metadata.
UserMetadata `protobuf:"bytes,3,opt,name=UserMetadata,proto3,embedded=UserMetadata" json:""`
+14
View File
@@ -88,6 +88,10 @@ type SAMLConnector interface {
GetEncryptionKeyPair() *AsymmetricKeyPair
// SetEncryptionKeyPair sets the key pair for SAML assertions.
SetEncryptionKeyPair(k *AsymmetricKeyPair)
// GetAllowIDPInitiated returns whether the identity provider can initiate a login or not.
GetAllowIDPInitiated() bool
// SetAllowIDPInitiated sets whether the identity provider can initiate a login or not.
SetAllowIDPInitiated(bool)
}
// NewSAMLConnector returns a new SAMLConnector based off a name and SAMLConnectorSpecV2.
@@ -332,6 +336,16 @@ func (o *SAMLConnectorV2) SetEncryptionKeyPair(k *AsymmetricKeyPair) {
o.Spec.EncryptionKeyPair = k
}
// GetAllowIDPInitiated returns whether the identity provider can initiate a login or not.
func (o *SAMLConnectorV2) GetAllowIDPInitiated() bool {
return o.Spec.AllowIDPInitiated
}
// SetAllowIDPInitiated sets whether the identity provider can initiate a login or not.
func (o *SAMLConnectorV2) SetAllowIDPInitiated(allow bool) {
o.Spec.AllowIDPInitiated = allow
}
// setStaticFields sets static resource header and metadata fields.
func (o *SAMLConnectorV2) setStaticFields() {
o.Kind = KindSAMLConnector
+944 -906
View File
File diff suppressed because it is too large Load Diff
+4
View File
@@ -2864,6 +2864,10 @@ message SAMLConnectorSpecV2 {
// EncryptionKeyPair is a key pair used for decrypting SAML assertions.
AsymmetricKeyPair EncryptionKeyPair = 13
[ (gogoproto.nullable) = true, (gogoproto.jsontag) = "assertion_key_pair,omitempty" ];
// AllowIDPInitiated is a flag that indicates if the connector can be used for IdP-initiated
// logins.
bool AllowIDPInitiated = 14
[ (gogoproto.nullable) = true, (gogoproto.jsontag) = "allow_idp_initiated,omitempty" ];
}
// SAMLAuthRequest is a request to authenticate with SAML
+1 -1
Submodule e updated: 393cd15422...7f65ada14e
+3 -2
View File
@@ -953,7 +953,8 @@ func (s *APIServer) createSAMLAuthRequest(auth ClientI, w http.ResponseWriter, r
}
type validateSAMLResponseReq struct {
Response string `json:"response"`
Response string `json:"response"`
ConnectorID string `json:"connector_id,omitempty"`
}
// samlAuthRawResponse is returned when auth server validated callback parameters
@@ -981,7 +982,7 @@ func (s *APIServer) validateSAMLResponse(auth ClientI, w http.ResponseWriter, r
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
response, err := auth.ValidateSAMLResponse(r.Context(), req.Response)
response, err := auth.ValidateSAMLResponse(r.Context(), req.Response, req.ConnectorID)
if err != nil {
return nil, trace.Wrap(err)
}
+6 -4
View File
@@ -20,7 +20,6 @@ limitations under the License.
// * Authority server itself that implements signing and acl logic
// * HTTP server wrapper for authority server
// * HTTP client wrapper
//
package auth
import (
@@ -166,6 +165,9 @@ func NewServer(cfg *InitConfig, opts ...ServerOption) (*Server, error) {
if cfg.Enforcer == nil {
cfg.Enforcer = local.NewNoopEnforcer()
}
if cfg.AssertionReplayService == nil {
cfg.AssertionReplayService = local.NewAssertionReplayService(cfg.Backend)
}
if cfg.KeyStoreConfig.RSAKeyPairSource == nil {
native.PrecomputeKeys()
cfg.KeyStoreConfig.RSAKeyPairSource = native.GenerateKeyPair
@@ -223,7 +225,7 @@ func NewServer(cfg *InitConfig, opts ...ServerOption) (*Server, error) {
closeCtx: closeCtx,
emitter: cfg.Emitter,
streamer: cfg.Streamer,
unstable: local.NewUnstableService(cfg.Backend),
unstable: local.NewUnstableService(cfg.Backend, cfg.AssertionReplayService),
Services: services,
Cache: services,
keyStore: keyStore,
@@ -338,8 +340,8 @@ var (
// Server keeps the cluster together. It acts as a certificate authority (CA) for
// a cluster and:
// - generates the keypair for the node it's running on
// - invites other SSH nodes to a cluster, by issuing invite tokens
// - adds other SSH nodes to a cluster, by checking their token and signing their keys
// - invites other SSH nodes to a cluster, by issuing invite tokens
// - adds other SSH nodes to a cluster, by checking their token and signing their keys
// - same for users and their sessions
// - checks public keys to see if they're signed by it (can be trusted or not)
type Server struct {
+2 -2
View File
@@ -2861,9 +2861,9 @@ func (a *ServerWithRoles) CreateSAMLAuthRequest(ctx context.Context, req types.S
}
// ValidateSAMLResponse validates SAML auth response.
func (a *ServerWithRoles) ValidateSAMLResponse(ctx context.Context, re string) (*SAMLAuthResponse, error) {
func (a *ServerWithRoles) ValidateSAMLResponse(ctx context.Context, re string, connectorID string) (*SAMLAuthResponse, error) {
// auth callback is it's own authz, no need to check extra permissions
return a.authServer.ValidateSAMLResponse(ctx, re)
return a.authServer.ValidateSAMLResponse(ctx, re, connectorID)
}
// GetSAMLAuthRequest returns SAML auth request if found.
+4 -3
View File
@@ -941,9 +941,10 @@ func (c *Client) ValidateOIDCAuthCallback(ctx context.Context, q url.Values) (*O
}
// ValidateSAMLResponse validates response returned by SAML identity provider
func (c *Client) ValidateSAMLResponse(ctx context.Context, re string) (*SAMLAuthResponse, error) {
func (c *Client) ValidateSAMLResponse(ctx context.Context, re string, connectorID string) (*SAMLAuthResponse, error) {
out, err := c.PostJSON(ctx, c.Endpoint("saml", "requests", "validate"), validateSAMLResponseReq{
Response: re,
Response: re,
ConnectorID: connectorID,
})
if err != nil {
return nil, trace.Wrap(err)
@@ -1426,7 +1427,7 @@ type IdentityService interface {
// CreateSAMLAuthRequest creates SAML AuthnRequest
CreateSAMLAuthRequest(ctx context.Context, req types.SAMLAuthRequest) (*types.SAMLAuthRequest, error)
// ValidateSAMLResponse validates SAML auth response
ValidateSAMLResponse(ctx context.Context, re string) (*SAMLAuthResponse, error)
ValidateSAMLResponse(ctx context.Context, re string, connectorID string) (*SAMLAuthResponse, error)
// GetSAMLAuthRequest returns SAML auth request if found
GetSAMLAuthRequest(ctx context.Context, authRequestID string) (*types.SAMLAuthRequest, error)
+4
View File
@@ -173,6 +173,7 @@ type InitConfig struct {
// WindowsServices is a service that manages Windows desktop resources.
WindowsDesktops services.WindowsDesktops
// SessionTrackerService is a service that manages trackers for all active sessions.
SessionTrackerService services.SessionTrackerService
// Enforcer is used to enforce Teleport Enterprise license compliance.
@@ -183,6 +184,9 @@ type InitConfig struct {
// TraceClient is used to forward spans to the upstream telemetry collector
TraceClient otlptrace.Client
// AssertionReplayService is a service that mitigatates SSO assertion replay.
*local.AssertionReplayService
}
// Init instantiates and configures an instance of AuthServer
+89 -38
View File
@@ -125,6 +125,19 @@ func (a *Server) CreateSAMLAuthRequest(ctx context.Context, req types.SAMLAuthRe
return &req, nil
}
func (a *Server) getSAMLConnectorAndProviderByID(ctx context.Context, connectorID string) (types.SAMLConnector, *saml2.SAMLServiceProvider, error) {
connector, err := a.Identity.GetSAMLConnector(ctx, connectorID, true)
if err != nil {
return nil, nil, trace.Wrap(err)
}
provider, err := a.getSAMLProvider(connector)
if err != nil {
return nil, nil, trace.Wrap(err)
}
return connector, provider, nil
}
func (a *Server) getSAMLConnectorAndProvider(ctx context.Context, req types.SAMLAuthRequest) (types.SAMLConnector, *saml2.SAMLServiceProvider, error) {
if req.SSOTestFlow {
if req.ConnectorSpec == nil {
@@ -157,16 +170,7 @@ func (a *Server) getSAMLConnectorAndProvider(ctx context.Context, req types.SAML
}
// regular execution flow
connector, err := a.GetSAMLConnector(ctx, req.ConnectorID, true)
if err != nil {
return nil, nil, trace.Wrap(err)
}
provider, err := a.getSAMLProvider(connector)
if err != nil {
return nil, nil, trace.Wrap(err)
}
return connector, provider, nil
return a.getSAMLConnectorAndProviderByID(ctx, req.ConnectorID)
}
func (a *Server) getSAMLProvider(conn types.SAMLConnector) (*saml2.SAMLServiceProvider, error) {
@@ -224,7 +228,12 @@ func (a *Server) calculateSAMLUser(diagCtx *ssoDiagContext, connector types.SAML
return nil, trace.Wrap(err)
}
roleTTL := roles.AdjustSessionTTL(apidefaults.MaxCertDuration)
p.sessionTTL = utils.MinTTL(roleTTL, request.CertTTL)
if request != nil {
p.sessionTTL = utils.MinTTL(roleTTL, request.CertTTL)
} else {
p.sessionTTL = roleTTL
}
return &p, nil
}
@@ -325,16 +334,12 @@ func ParseSAMLInResponseTo(response string) (string, error) {
return "", trace.BadParameter("unable to parse response")
}
// teleport only supports sending party initiated flows (Teleport sends an
// AuthnRequest to the IdP and gets a SAMLResponse from the IdP). identity
// provider initiated flows (where Teleport gets an unsolicited SAMLResponse
// from the IdP) are not supported.
// 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 {
message := "teleport does not support initiating login from a SAML identity provider, login must be initiated from either the Teleport Web UI or CLI"
log.Infof(message)
return "", trace.NotImplemented(message)
return "", trace.NotFound("missing InResponseTo attribute")
}
if responseTo.Value == "" {
return "", trace.BadParameter("InResponseTo can not be empty")
@@ -363,7 +368,7 @@ type SAMLAuthResponse struct {
}
// ValidateSAMLResponse consumes attribute statements from SAML identity provider
func (a *Server) ValidateSAMLResponse(ctx context.Context, samlResponse string) (*SAMLAuthResponse, error) {
func (a *Server) ValidateSAMLResponse(ctx context.Context, samlResponse string, connectorID string) (*SAMLAuthResponse, error) {
event := &apievents.UserLogin{
Metadata: apievents.Metadata{
Type: events.UserLoginEvent,
@@ -373,7 +378,7 @@ func (a *Server) ValidateSAMLResponse(ctx context.Context, samlResponse string)
diagCtx := a.newSSODiagContext(types.KindSAML)
auth, err := a.validateSAMLResponse(ctx, diagCtx, samlResponse)
auth, err := a.validateSAMLResponse(ctx, diagCtx, samlResponse, connectorID)
diagCtx.info.Error = trace.UserMessage(err)
diagCtx.writeToBackend(ctx)
@@ -417,22 +422,51 @@ func (a *Server) ValidateSAMLResponse(ctx context.Context, samlResponse string)
return auth, nil
}
func (a *Server) validateSAMLResponse(ctx context.Context, diagCtx *ssoDiagContext, samlResponse string) (*SAMLAuthResponse, error) {
func (a *Server) 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 := a.unstable.RecognizeSSOAssertion(ctx, connector.GetName(), assertion.SessionIndex, assertion.NameID, *assertion.SessionNotOnOrAfter)
return trace.Wrap(err)
}
func (a *Server) 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)
if err != nil {
return nil, trace.Wrap(err)
}
diagCtx.requestID = requestID
switch {
case trace.IsNotFound(err):
if connectorID == "" {
return nil, trace.BadParameter("ACS URI did not include a valid SAML connector ID parameter")
}
request, err := a.GetSAMLAuthRequest(ctx, requestID)
if err != nil {
return nil, trace.Wrap(err, "Failed to get SAML Auth Request")
}
diagCtx.info.TestFlow = request.SSOTestFlow
idpInitiated = true
connector, provider, err = a.getSAMLConnectorAndProviderByID(ctx, connectorID)
if err != nil {
return nil, trace.Wrap(err, "Failed to get SAML connector and provider")
}
case err != nil:
trace.Wrap(err)
default:
diagCtx.requestID = requestID
request, err = a.Identity.GetSAMLAuthRequest(ctx, requestID)
if err != nil {
return nil, trace.Wrap(err, "Failed to get SAML Auth Request")
}
connector, provider, err := a.getSAMLConnectorAndProvider(ctx, *request)
if err != nil {
return nil, trace.Wrap(err, "Failed to get SAML connector and provider")
diagCtx.info.TestFlow = request.SSOTestFlow
connector, provider, err = a.getSAMLConnectorAndProvider(ctx, *request)
if err != nil {
return nil, trace.Wrap(err, "Failed to get SAML connector and provider")
}
}
assertionInfo, err := provider.RetrieveAssertionInfo(samlResponse)
@@ -444,6 +478,16 @@ func (a *Server) validateSAMLResponse(ctx context.Context, diagCtx *ssoDiagConte
diagCtx.info.SAMLAssertionInfo = (*types.AssertionInfo)(assertionInfo)
}
if idpInitiated {
if err := a.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.")
}
@@ -492,14 +536,13 @@ func (a *Server) validateSAMLResponse(ctx context.Context, diagCtx *ssoDiagConte
SessionTTL: types.Duration(params.sessionTTL),
}
user, err := a.createSAMLUser(params, request.SSOTestFlow)
user, err := a.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{
Req: *request,
Identity: types.ExternalIdentity{
ConnectorID: params.connectorName,
Username: params.username,
@@ -507,14 +550,22 @@ func (a *Server) validateSAMLResponse(ctx context.Context, diagCtx *ssoDiagConte
Username: user.GetName(),
}
if request != nil {
auth.Req = *request
} else {
auth.Req = types.SAMLAuthRequest{
CreateWebSession: true,
}
}
// In test flow skip signing and creating web sessions.
if request.SSOTestFlow {
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.CreateWebSession {
if request == nil || request.CreateWebSession {
session, err := a.createWebSession(ctx, types.NewWebSessionRequest{
User: user.GetName(),
Roles: user.GetRoles(),
@@ -530,7 +581,7 @@ func (a *Server) validateSAMLResponse(ctx context.Context, diagCtx *ssoDiagConte
}
// If a public key was provided, sign it and return a certificate.
if len(request.PublicKey) != 0 {
if request != nil && len(request.PublicKey) != 0 {
sshCert, tlsCert, err := a.createSessionCert(user, params.sessionTTL, request.PublicKey, request.Compatibility, request.RouteToCluster, request.KubernetesCluster)
if err != nil {
return nil, trace.Wrap(err, "Failed to create session certificate.")
+3 -3
View File
@@ -399,7 +399,7 @@ func TestServer_ValidateSAMLResponse(t *testing.T) {
a.SetClock(clock)
// empty response gives error.
response, err := a.ValidateSAMLResponse(context.Background(), "")
response, err := a.ValidateSAMLResponse(context.Background(), "", "")
require.Nil(t, response)
require.Error(t, err)
@@ -520,13 +520,13 @@ V115UGOwvjOOxmOFbYBn865SHgMndFtr</ds:X509Certificate></ds:X509Data></ds:KeyInfo>
require.NoError(t, err)
// check ValidateSAMLResponse
response, err = a.ValidateSAMLResponse(context.Background(), base64.StdEncoding.EncodeToString([]byte(respOkta)))
response, err = a.ValidateSAMLResponse(context.Background(), base64.StdEncoding.EncodeToString([]byte(respOkta)), "")
require.NoError(t, err)
require.NotNil(t, response)
// check internal method, validate diagnostic outputs.
diagCtx := a.newSSODiagContext(types.KindSAML)
auth, err := a.validateSAMLResponse(context.Background(), diagCtx, base64.StdEncoding.EncodeToString([]byte(respOkta)))
auth, err := a.validateSAMLResponse(context.Background(), diagCtx, base64.StdEncoding.EncodeToString([]byte(respOkta)), "")
require.NoError(t, err)
// ensure diag info got stored and is identical.
+11 -3
View File
@@ -369,13 +369,21 @@ func (l *Backend) Create(ctx context.Context, i backend.Item) (*backend.Lease, e
return trace.Wrap(err)
}
}
stmt, err := tx.PrepareContext(ctx, "INSERT INTO kv(key, modified, expires, value) values(?, ?, ?, ?)")
rows, err := tx.QueryContext(ctx, "SELECT key, value, expires, modified FROM kv WHERE key = ? AND expires <= ? LIMIT 1", string(i.Key), created)
if err != nil {
return trace.Wrap(err)
}
defer stmt.Close()
defer rows.Close()
if _, err := stmt.ExecContext(ctx, string(i.Key), id(created), expires(i.Expires), i.Value); err != nil {
if rows.Next() {
err = l.deleteInTransaction(ctx, i.Key, tx)
if err != nil {
return trace.Wrap(err)
}
}
if _, err := tx.ExecContext(ctx, "INSERT INTO kv(key, modified, expires, value) values(?, ?, ?, ?)", string(i.Key), id(created), expires(i.Expires), i.Value); err != nil {
return trace.Wrap(err)
}
return nil
+57
View File
@@ -0,0 +1,57 @@
/*
Copyright 2022 Gravitational, Inc.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
package local
import (
"context"
"time"
"github.com/gravitational/teleport/lib/backend"
"github.com/gravitational/trace"
)
const assertionReplayPrefix = "recognized_assertions"
// AssertionReplayService tracks used SSO assertions to mitigate replay attacks.
// Assertions are automatically derecognized when their signed expiry passes.
type AssertionReplayService struct {
bk backend.Backend
}
// NewAssertionReplayService creates a new instance of AssertionReplayService.
func NewAssertionReplayService(bk backend.Backend) *AssertionReplayService {
return &AssertionReplayService{bk: bk}
}
// RecognizeSSOAssertion will remember a new assertion until it becomes invalid.
// This will error with `trace.AlreadyExists` if the assertion has been previously recognized.
//
// `safeAfter` must be either at or after the point in time that a given SSO assertion becomes invalid in order to mitigate replay attacks.
// This function shouldn't be used if the assertion never verifiably expires.
func (s *AssertionReplayService) RecognizeSSOAssertion(ctx context.Context, connectorID string, assertionID string, user string, safeAfter time.Time) error {
key := backend.Key(assertionReplayPrefix, connectorID, assertionID)
item := backend.Item{Key: key, Value: []byte(user), Expires: safeAfter}
_, err := s.bk.Create(ctx, item)
switch {
case trace.IsAlreadyExists(err):
return trace.AlreadyExists("Assertion %q already recognized for user %v", assertionID, user)
case err != nil:
return trace.Wrap(err)
default:
return nil
}
}
@@ -0,0 +1,57 @@
/*
Copyright 2022 Gravitational, Inc.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
package local
import (
"context"
"testing"
"time"
"github.com/google/uuid"
"github.com/gravitational/teleport/lib/backend/memory"
"github.com/stretchr/testify/require"
)
func TestAssertionReplayService(t *testing.T) {
t.Parallel()
ctx := context.Background()
delay := func(t time.Duration) time.Time { return time.Now().UTC().Add(t) }
bk, err := memory.New(memory.Config{})
require.NoError(t, err)
service := NewAssertionReplayService(bk)
id := make([]string, 2)
for i := range id {
id[i] = uuid.New().String()
}
// first time foo
require.NoError(t, service.RecognizeSSOAssertion(ctx, "", id[0], "foo", delay(time.Hour)))
// second time foo
require.Error(t, service.RecognizeSSOAssertion(ctx, "", id[0], "foo", delay(time.Hour)))
// first time bar
require.NoError(t, service.RecognizeSSOAssertion(ctx, "", id[1], "bar", delay(time.Millisecond)))
time.Sleep(time.Second)
// assertion has expired, no risk of replay
require.NoError(t, service.RecognizeSSOAssertion(ctx, "", id[1], "bar", delay(time.Hour)))
// assertion should still exist
require.Error(t, service.RecognizeSSOAssertion(ctx, "", id[1], "bar", delay(time.Hour)))
}
+3 -2
View File
@@ -33,11 +33,12 @@ const assertionTTL = time.Minute * 10
// that don't fit into, or merit the change of, one of the primary service interfaces.
type UnstableService struct {
backend.Backend
*AssertionReplayService
}
// NewUnstableService returns new unstable service instance.
func NewUnstableService(backend backend.Backend) UnstableService {
return UnstableService{Backend: backend}
func NewUnstableService(backend backend.Backend, assertion *AssertionReplayService) UnstableService {
return UnstableService{backend, assertion}
}
func (s UnstableService) AssertSystemRole(ctx context.Context, req proto.UnstableSystemRoleAssertion) error {
+2 -1
View File
@@ -44,7 +44,8 @@ func TestSystemRoleAssertions(t *testing.T) {
defer backend.Close()
unstable := NewUnstableService(backend)
assertion := NewAssertionReplayService(backend)
unstable := NewUnstableService(backend, assertion)
_, err = unstable.GetSystemRoleAssertions(ctx, serverID, assertionID)
require.True(t, trace.IsNotFound(err))
+8 -5
View File
@@ -548,6 +548,7 @@ func (h *Handler) bindDefaultEndpoints(challengeLimiter *limiter.RateLimiter) {
// SAML 2.0 handlers
h.POST("/webapi/saml/acs", h.WithRedirect(h.samlACS))
h.POST("/webapi/saml/acs/:connector", h.WithRedirect(h.samlACS))
h.GET("/webapi/saml/sso", h.WithMetaRedirect(h.samlSSO))
h.POST("/webapi/saml/login/console", httplib.MakeHandler(h.samlSSOConsole))
@@ -1286,7 +1287,7 @@ func (h *Handler) githubCallback(w http.ResponseWriter, r *http.Request, p httpr
clientRedirectURL: response.Req.ClientRedirectURL,
}
if err := ssoSetWebSessionAndRedirectURL(w, r, res); err != nil {
if err := ssoSetWebSessionAndRedirectURL(w, r, res, true); err != nil {
logger.WithError(err).Error("Error setting web session.")
return client.LoginFailedRedirectURL
}
@@ -1392,7 +1393,7 @@ func (h *Handler) oidcCallback(w http.ResponseWriter, r *http.Request, p httprou
clientRedirectURL: response.Req.ClientRedirectURL,
}
if err := ssoSetWebSessionAndRedirectURL(w, r, res); err != nil {
if err := ssoSetWebSessionAndRedirectURL(w, r, res, true); err != nil {
logger.WithError(err).Error("Error setting web session.")
return client.LoginFailedRedirectURL
}
@@ -3061,9 +3062,11 @@ type ssoCallbackResponse struct {
clientRedirectURL string
}
func ssoSetWebSessionAndRedirectURL(w http.ResponseWriter, r *http.Request, response *ssoCallbackResponse) error {
if err := csrf.VerifyToken(response.csrfToken, r); err != nil {
return trace.Wrap(err)
func ssoSetWebSessionAndRedirectURL(w http.ResponseWriter, r *http.Request, response *ssoCallbackResponse, verifyCSRF bool) error {
if verifyCSRF {
if err := csrf.VerifyToken(response.csrfToken, r); err != nil {
return trace.Wrap(err)
}
}
if err := SetSessionCookie(w, response.username, response.sessionName); err != nil {
+8 -3
View File
@@ -97,7 +97,7 @@ func (h *Handler) samlACS(w http.ResponseWriter, r *http.Request, p httprouter.P
return client.LoginFailedRedirectURL
}
response, err := h.cfg.ProxyClient.ValidateSAMLResponse(r.Context(), samlResponse)
response, err := h.cfg.ProxyClient.ValidateSAMLResponse(r.Context(), samlResponse, p.ByName("connector"))
if err != nil {
logger.WithError(err).Error("Error while processing callback.")
@@ -125,14 +125,19 @@ func (h *Handler) samlACS(w http.ResponseWriter, r *http.Request, p httprouter.P
if response.Req.CreateWebSession {
logger.Debug("Redirecting to web browser.")
redirect := response.Req.ClientRedirectURL
if redirect == "" {
redirect = "/web/"
}
res := &ssoCallbackResponse{
csrfToken: response.Req.CSRFToken,
username: response.Username,
sessionName: response.Session.GetName(),
clientRedirectURL: response.Req.ClientRedirectURL,
clientRedirectURL: redirect,
}
if err := ssoSetWebSessionAndRedirectURL(w, r, res); err != nil {
if err := ssoSetWebSessionAndRedirectURL(w, r, res, response.Req.CSRFToken != ""); err != nil {
logger.WithError(err).Error("Error setting web session.")
return client.LoginFailedRedirectURL
}