diff --git a/pkg/auth/saml_cleanup.go b/pkg/auth/saml_cleanup.go index d60520dd6..3a2d23a0c 100644 --- a/pkg/auth/saml_cleanup.go +++ b/pkg/auth/saml_cleanup.go @@ -74,7 +74,7 @@ func (c *Cleaner) Run(ctx context.Context) error { } func (c *Cleaner) cleanup(ctx context.Context) error { - var assertionsDeleted, requestsDeleted, relayStatesDeleted int64 + var assertionsDeleted, requestsDeleted int64 err := c.pg.WithConn( ctx, @@ -91,12 +91,6 @@ func (c *Cleaner) cleanup(ctx context.Context) error { } requestsDeleted = count - count, err = CleanupExpiredRelayStates(ctx, conn) - if err != nil { - return err - } - relayStatesDeleted = count - return nil }, ) @@ -105,11 +99,10 @@ func (c *Cleaner) cleanup(ctx context.Context) error { return err } - if assertionsDeleted > 0 || requestsDeleted > 0 || relayStatesDeleted > 0 { + if assertionsDeleted > 0 || requestsDeleted > 0 { c.logger.InfoCtx(ctx, "cleaned up expired SAML data", log.Int64("assertions", assertionsDeleted), - log.Int64("requests", requestsDeleted), - log.Int64("relay_states", relayStatesDeleted)) + log.Int64("requests", requestsDeleted)) } return nil diff --git a/pkg/auth/saml_service.go b/pkg/auth/saml_service.go index 992401058..c19685596 100644 --- a/pkg/auth/saml_service.go +++ b/pkg/auth/saml_service.go @@ -20,7 +20,6 @@ import ( "crypto/x509" "encoding/base64" "encoding/pem" - "encoding/xml" "errors" "fmt" "net/http" @@ -232,27 +231,6 @@ func (s *SAMLService) GetAcsURL() string { return fmt.Sprintf("%s/connect/saml/consume", s.baseURL) } -func parseRawSAMLResponse(encodedResponse string) (*saml.Assertion, error) { - rawResponseBuf, err := base64.StdEncoding.DecodeString(encodedResponse) - if err != nil { - return nil, fmt.Errorf("cannot decode base64: %w", err) - } - - var response saml.Response - if err := xml.Unmarshal(rawResponseBuf, &response); err != nil { - return nil, fmt.Errorf("cannot unmarshal response: %w", err) - } - - if response.Assertion == nil { - if response.EncryptedAssertion != nil { - return nil, fmt.Errorf("response contains encrypted assertion which cannot be parsed without SP private key") - } - return nil, fmt.Errorf("response contains no assertion") - } - - return response.Assertion, nil -} - func (s *SAMLService) GetServiceProvider( ctx context.Context, config *coredata.SAMLConfiguration, @@ -271,18 +249,24 @@ func (s *SAMLService) GetServiceProvider( return nil, ErrInvalidURL{Field: "ACS", URL: s.GetAcsURL(), Err: err} } + metadataURL, err := url.Parse(s.GetEntityID()) + if err != nil { + return nil, ErrInvalidURL{Field: "Metadata", URL: s.GetEntityID(), Err: err} + } + idpSSOURL, err := url.Parse(config.IdPSsoURL) if err != nil { return nil, ErrInvalidURL{Field: "IdP SSO", URL: config.IdPSsoURL, Err: err} } sp := &saml.ServiceProvider{ - EntityID: s.GetEntityID(), - Key: s.privateKey, - Certificate: s.certificate, - MetadataURL: *acsURL, - AcsURL: *acsURL, - SloURL: *acsURL, + EntityID: s.GetEntityID(), + Key: s.privateKey, + Certificate: s.certificate, + MetadataURL: *metadataURL, + AcsURL: *acsURL, + SloURL: *acsURL, + AllowIDPInitiated: true, IDPMetadata: &saml.EntityDescriptor{ EntityID: config.IdPEntityID, IDPSSODescriptors: []saml.IDPSSODescriptor{ @@ -355,14 +339,8 @@ func (s *SAMLService) InitiateSAMLLogin( return "", ErrCannotCreateAuthRequest{Err: err} } - relayStateToken, err := coredata.GenerateSecureToken() - if err != nil { - return "", fmt.Errorf("cannot generate relay state token: %w", err) - } - now := time.Now() requestExpiry := now.Add(10 * time.Minute) - relayStateExpiry := now.Add(15 * time.Minute) err = s.pg.WithTx( ctx, @@ -373,20 +351,9 @@ func (s *SAMLService) InitiateSAMLLogin( CreatedAt: now, ExpiresAt: requestExpiry, } - if err := samlRequest.Insert(ctx, tx, scope); err != nil { - return fmt.Errorf("cannot store SAML request: %w", err) - } - relayState := coredata.SAMLRelayState{ - Token: relayStateToken, - OrganizationID: organizationID, - SAMLConfigID: config.ID, - RequestID: authReq.ID, - CreatedAt: now, - ExpiresAt: relayStateExpiry, - } - if err := relayState.Insert(ctx, tx, scope); err != nil { - return fmt.Errorf("cannot store relay state: %w", err) + if err := samlRequest.Insert(ctx, tx, scope); err != nil { + return fmt.Errorf("cannot insert SAML request: %w", err) } return nil @@ -396,7 +363,7 @@ func (s *SAMLService) InitiateSAMLLogin( return "", err } - redirectURL, err := authReq.Redirect(relayStateToken, sp) + redirectURL, err := authReq.Redirect(config.ID.String(), sp) if err != nil { return "", ErrCannotGenerateRedirectURL{Err: err} } @@ -413,85 +380,36 @@ type SAMLUserInfo struct { SAMLConfigID gid.GID } -func (s *SAMLService) loadContextForSPInitiated( +func (s *SAMLService) loadConfigFromRelayState( ctx context.Context, - relayStateToken string, - now time.Time, -) (*coredata.SAMLConfiguration, *coredata.Organization, string, error) { - var relayState coredata.SAMLRelayState - var samlRequest coredata.SAMLRequest - var org coredata.Organization - var config coredata.SAMLConfiguration - - err := s.pg.WithTx(ctx, func(tx pg.Conn) error { - if err := relayState.Load(ctx, tx, relayStateToken); err != nil { - return fmt.Errorf("invalid relay state: %w", err) - } - - if relayState.IsExpired(now) { - return coredata.ErrRelayStateExpired{Token: relayStateToken, ExpiresAt: relayState.ExpiresAt} - } - - if err := samlRequest.Load(ctx, tx, relayState.RequestID, relayState.OrganizationID); err != nil { - return fmt.Errorf("invalid SAML request: %w", err) - } - - if samlRequest.IsExpired(now) { - return coredata.ErrSAMLRequestExpired{RequestID: relayState.RequestID, ExpiresAt: samlRequest.ExpiresAt} - } - - if err := org.LoadByID(ctx, tx, coredata.NewNoScope(), relayState.OrganizationID); err != nil { - return fmt.Errorf("organization not found: %w", err) - } - - scope := coredata.NewScope(org.TenantID) - if err := config.LoadByID(ctx, tx, scope, relayState.SAMLConfigID); err != nil { - return fmt.Errorf("cannot load SAML configuration: %w", err) - } - - if err := relayState.Delete(ctx, tx); err != nil { - return fmt.Errorf("cannot delete relay state: %w", err) - } - if err := samlRequest.Delete(ctx, tx); err != nil { - return fmt.Errorf("cannot delete SAML request: %w", err) - } - - return nil - }) - if err != nil { - return nil, nil, "", err - } - - return &config, &org, samlRequest.ID, nil -} - -func (s *SAMLService) loadContextForIDPInitiated( - ctx context.Context, - samlConfigIDParam string, + relayStateValue string, ) (*coredata.SAMLConfiguration, *coredata.Organization, error) { - if samlConfigIDParam == "" { - return nil, nil, fmt.Errorf("IDP-initiated login requires 'c' query parameter with SAML config ID") + if relayStateValue == "" { + return nil, nil, fmt.Errorf("RelayState is required and must contain SAML config ID") } - samlConfigID, err := gid.ParseGID(samlConfigIDParam) + samlConfigID, err := gid.ParseGID(relayStateValue) if err != nil { - return nil, nil, fmt.Errorf("invalid 'c' parameter: %w", err) + return nil, nil, fmt.Errorf("invalid SAML config ID in RelayState: %w", err) } var config coredata.SAMLConfiguration var org coredata.Organization - err = s.pg.WithConn(ctx, func(conn pg.Conn) error { - if err := config.LoadByID(ctx, conn, coredata.NewNoScope(), samlConfigID); err != nil { - return fmt.Errorf("cannot load SAML configuration: %w", err) - } + err = s.pg.WithConn( + ctx, + func(conn pg.Conn) error { + if err := config.LoadByID(ctx, conn, coredata.NewNoScope(), samlConfigID); err != nil { + return fmt.Errorf("cannot load SAML configuration: %w", err) + } - if err := org.LoadByID(ctx, conn, coredata.NewNoScope(), config.OrganizationID); err != nil { - return fmt.Errorf("organization not found: %w", err) - } + if err := org.LoadByID(ctx, conn, coredata.NewNoScope(), config.OrganizationID); err != nil { + return fmt.Errorf("organization not found: %w", err) + } - return nil - }) + return nil + }, + ) if err != nil { return nil, nil, err } @@ -508,28 +426,10 @@ func (s *SAMLService) HandleSAMLAssertion( return nil, fmt.Errorf("missing SAMLResponse in request") } - relayStateToken := req.FormValue("RelayState") - samlConfigIDParam := req.URL.Query().Get("c") - now := time.Now() - - var config *coredata.SAMLConfiguration - var org *coredata.Organization - var possibleRequestIDs []string - var err error - - if relayStateToken != "" { - var requestID string - config, org, requestID, err = s.loadContextForSPInitiated(ctx, relayStateToken, now) - if err != nil { - return nil, err - } - possibleRequestIDs = []string{requestID} - } else { - config, org, err = s.loadContextForIDPInitiated(ctx, samlConfigIDParam) - if err != nil { - return nil, err - } - possibleRequestIDs = []string{} + relayStateValue := req.FormValue("RelayState") + config, org, err := s.loadConfigFromRelayState(ctx, relayStateValue) + if err != nil { + return nil, err } if !config.Enabled { @@ -548,14 +448,28 @@ func (s *SAMLService) HandleSAMLAssertion( req.URL.Host = req.Host } + now := time.Now() + + var possibleRequestIDs []string + err = s.pg.WithConn( + ctx, + func(conn pg.Conn) error { + requestIDs, err := coredata.LoadValidRequestIDsForOrganization(ctx, conn, config.OrganizationID, now) + if err != nil { + return err + } + possibleRequestIDs = requestIDs + + return nil + }, + ) + if err != nil { + return nil, fmt.Errorf("cannot load valid request IDs: %w", err) + } + assertion, err := sp.ParseResponse(req, possibleRequestIDs) if err != nil { - return nil, fmt.Errorf( - "cannot parse SAML response (SP EntityID: %s, IdP EntityID: %s): %w", - s.GetEntityID(), - config.IdPEntityID, - err, - ) + return nil, fmt.Errorf("cannot parse SAML response: %w", err) } if err := ValidateAssertion(assertion, s.GetEntityID(), now); err != nil { @@ -571,14 +485,22 @@ func (s *SAMLService) HandleSAMLAssertion( } scope := coredata.NewScope(org.TenantID) - err = s.pg.WithTx(ctx, func(tx pg.Conn) error { - return PreventReplayAttack(ctx, tx, scope, assertion.ID, config.OrganizationID, expiresAt) - }) + err = s.pg.WithTx( + ctx, + func(tx pg.Conn) error { + if err := PreventReplayAttack(ctx, tx, scope, assertion.ID, config.OrganizationID, expiresAt); err != nil { + return fmt.Errorf("cannot prevent replay attack: %w", err) + } + + return nil + }, + ) if err != nil { var replayAttackErr *coredata.ErrAssertionAlreadyUsed if errors.As(err, &replayAttackErr) { return nil, ErrReplayAttackDetected{AssertionID: assertion.ID, Err: replayAttackErr} } + return nil, fmt.Errorf("cannot prevent replay attack: %w", err) } } @@ -598,6 +520,7 @@ func (s *SAMLService) HandleSAMLAssertion( if err != nil { return nil, ErrCannotExtractUserAttributes{Err: fmt.Errorf("cannot extract domain from email: %w", err)} } + if actualEmailDomain != config.EmailDomain { return nil, fmt.Errorf("email domain mismatch: assertion contains email with domain %s but SAML config is for domain %s", actualEmailDomain, config.EmailDomain) } diff --git a/pkg/auth/saml_validator.go b/pkg/auth/saml_validator.go index 17692404e..95badcb5e 100644 --- a/pkg/auth/saml_validator.go +++ b/pkg/auth/saml_validator.go @@ -101,7 +101,3 @@ func CleanupExpiredAssertions(ctx context.Context, conn pg.Conn) (int64, error) func CleanupExpiredRequests(ctx context.Context, conn pg.Conn) (int64, error) { return coredata.DeleteExpiredSAMLRequests(ctx, conn, time.Now()) } - -func CleanupExpiredRelayStates(ctx context.Context, conn pg.Conn) (int64, error) { - return coredata.DeleteExpiredSAMLRelayStates(ctx, conn, time.Now()) -} diff --git a/pkg/coredata/migrations/20251114T000000Z.sql b/pkg/coredata/migrations/20251114T000000Z.sql new file mode 100644 index 000000000..8d92f853b --- /dev/null +++ b/pkg/coredata/migrations/20251114T000000Z.sql @@ -0,0 +1 @@ +DROP TABLE auth_saml_relay_states; diff --git a/pkg/coredata/saml_relay_state.go b/pkg/coredata/saml_relay_state.go deleted file mode 100644 index 0d23878ed..000000000 --- a/pkg/coredata/saml_relay_state.go +++ /dev/null @@ -1,156 +0,0 @@ -// Copyright (c) 2025 Probo Inc . -// -// Permission to use, copy, modify, and/or distribute this software for any -// purpose with or without fee is hereby granted, provided that the above -// copyright notice and this permission notice appear in all copies. -// -// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH -// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY -// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, -// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM -// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR -// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR -// PERFORMANCE OF THIS SOFTWARE. - -package coredata - -import ( - "context" - "crypto/rand" - "encoding/base64" - "fmt" - "time" - - "go.probo.inc/probo/pkg/gid" - "github.com/jackc/pgx/v5" - "go.gearno.de/kit/pg" -) - -type SAMLRelayState struct { - Token string `db:"token"` - OrganizationID gid.GID `db:"organization_id"` - SAMLConfigID gid.GID `db:"saml_config_id"` - RequestID string `db:"request_id"` - CreatedAt time.Time `db:"created_at"` - ExpiresAt time.Time `db:"expires_at"` -} - -type ErrRelayStateNotFound struct { - Token string -} - -func (e ErrRelayStateNotFound) Error() string { - return "relay state token not found or invalid" -} - -type ErrRelayStateExpired struct { - Token string - ExpiresAt time.Time -} - -func (e ErrRelayStateExpired) Error() string { - return fmt.Sprintf("relay state token expired at %v", e.ExpiresAt) -} - -func GenerateSecureToken() (string, error) { - b := make([]byte, 32) - _, err := rand.Read(b) - if err != nil { - return "", fmt.Errorf("cannot generate random token: %w", err) - } - - token := base64.URLEncoding.EncodeToString(b) - return token, nil -} - -func (s *SAMLRelayState) Insert( - ctx context.Context, - conn pg.Conn, - scope Scoper, -) error { - query := ` -INSERT INTO auth_saml_relay_states (token, tenant_id, organization_id, saml_config_id, request_id, created_at, expires_at) -VALUES (@token, @tenant_id, @organization_id, @saml_config_id, @request_id, @created_at, @expires_at) -` - - args := pgx.NamedArgs{ - "token": s.Token, - "tenant_id": scope.GetTenantID(), - "organization_id": s.OrganizationID, - "saml_config_id": s.SAMLConfigID, - "request_id": s.RequestID, - "created_at": s.CreatedAt, - "expires_at": s.ExpiresAt, - } - - _, err := conn.Exec(ctx, query, args) - if err != nil { - return fmt.Errorf("cannot insert saml_relay_state: %w", err) - } - - return nil -} - -func (s *SAMLRelayState) Load( - ctx context.Context, - conn pg.Conn, - token string, -) error { - query := ` -SELECT token, organization_id, saml_config_id, request_id, created_at, expires_at -FROM auth_saml_relay_states -WHERE token = @token -LIMIT 1 -` - - rows, err := conn.Query(ctx, query, pgx.NamedArgs{"token": token}) - if err != nil { - return fmt.Errorf("cannot query saml_relay_states: %w", err) - } - - state, err := pgx.CollectOneRow(rows, pgx.RowToStructByName[SAMLRelayState]) - if err == pgx.ErrNoRows { - return ErrRelayStateNotFound{Token: token} - } - if err != nil { - return fmt.Errorf("cannot collect saml_relay_state: %w", err) - } - - *s = state - return nil -} - -func (s *SAMLRelayState) IsExpired(now time.Time) bool { - return now.After(s.ExpiresAt) || now.Equal(s.ExpiresAt) -} - -func (s *SAMLRelayState) Delete( - ctx context.Context, - conn pg.Conn, -) error { - query := ` -DELETE FROM auth_saml_relay_states -WHERE token = @token -` - - _, err := conn.Exec(ctx, query, pgx.NamedArgs{"token": s.Token}) - if err != nil { - return fmt.Errorf("cannot delete saml_relay_state: %w", err) - } - - return nil -} - -func DeleteExpiredSAMLRelayStates(ctx context.Context, conn pg.Conn, now time.Time) (int64, error) { - query := ` -DELETE FROM auth_saml_relay_states -WHERE expires_at < @now -` - - result, err := conn.Exec(ctx, query, pgx.NamedArgs{"now": now}) - if err != nil { - return 0, fmt.Errorf("cannot delete expired saml_relay_states: %w", err) - } - - return result.RowsAffected(), nil -} diff --git a/pkg/coredata/saml_request.go b/pkg/coredata/saml_request.go index c206e72de..d3ba1ee01 100644 --- a/pkg/coredata/saml_request.go +++ b/pkg/coredata/saml_request.go @@ -130,6 +130,40 @@ WHERE id = @id return nil } +func LoadValidRequestIDsForOrganization( + ctx context.Context, + conn pg.Conn, + organizationID gid.GID, + now time.Time, +) ([]string, error) { + query := ` +SELECT id +FROM auth_saml_requests +WHERE organization_id = @organization_id AND expires_at > @now +` + + args := pgx.NamedArgs{ + "organization_id": organizationID, + "now": now, + } + + rows, err := conn.Query(ctx, query, args) + if err != nil { + return nil, fmt.Errorf("cannot query saml_requests: %w", err) + } + + requestIDs, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (string, error) { + var id string + err := row.Scan(&id) + return id, err + }) + if err != nil { + return nil, fmt.Errorf("cannot collect request IDs: %w", err) + } + + return requestIDs, nil +} + func DeleteExpiredSAMLRequests(ctx context.Context, conn pg.Conn, now time.Time) (int64, error) { query := ` DELETE FROM auth_saml_requests diff --git a/pkg/server/auth/saml_acs_handler.go b/pkg/server/auth/saml_acs_handler.go index 73518ece5..a6721d7c1 100644 --- a/pkg/server/auth/saml_acs_handler.go +++ b/pkg/server/auth/saml_acs_handler.go @@ -60,15 +60,6 @@ func SAMLACSHandler(samlSvc *authsvc.SAMLService, authSvc *authsvc.Service, auth return } - relayState := r.FormValue("RelayState") - samlConfigID := r.URL.Query().Get("c") - - if relayState != "" { - logger.InfoCtx(ctx, "processing SP-initiated SAML login") - } else { - logger.InfoCtx(ctx, "processing IDP-initiated SAML login", log.String("config_id", samlConfigID)) - } - userInfo, err := samlSvc.HandleSAMLAssertion(ctx, r) if err != nil { logger.ErrorCtx(ctx, "SAML authentication failed", log.Error(err))