mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
Add support for IdP-Initiated SAML2 login (#13924)
This commit is contained in:
@@ -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:""`
|
||||
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
|
||||
@@ -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
@@ -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.")
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)))
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user