Add IDP initiated flow support

Signed-off-by: Bryan Frimin <bryan@getprobo.com>
This commit is contained in:
Bryan Frimin
2025-11-12 09:06:02 +01:00
parent 5b6016c817
commit edce70141d
2 changed files with 123 additions and 82 deletions

View File

@@ -28,11 +28,11 @@ import (
"time" "time"
"github.com/crewjam/saml" "github.com/crewjam/saml"
"go.gearno.de/kit/log"
"go.gearno.de/kit/pg"
"go.probo.inc/probo/pkg/coredata" "go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/crypto/cipher" "go.probo.inc/probo/pkg/crypto/cipher"
"go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/gid"
"go.gearno.de/kit/log"
"go.gearno.de/kit/pg"
) )
type ( type (
@@ -413,84 +413,130 @@ type SAMLUserInfo struct {
SAMLConfigID gid.GID SAMLConfigID gid.GID
} }
func (s *SAMLService) loadContextForSPInitiated(
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,
) (*coredata.SAMLConfiguration, *coredata.Organization, error) {
if samlConfigIDParam == "" {
return nil, nil, fmt.Errorf("IDP-initiated login requires 'c' query parameter with SAML config ID")
}
samlConfigID, err := gid.ParseGID(samlConfigIDParam)
if err != nil {
return nil, nil, fmt.Errorf("invalid 'c' parameter: %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)
}
if err := org.LoadByID(ctx, conn, coredata.NewNoScope(), config.OrganizationID); err != nil {
return fmt.Errorf("organization not found: %w", err)
}
return nil
})
if err != nil {
return nil, nil, err
}
return &config, &org, nil
}
func (s *SAMLService) HandleSAMLAssertion( func (s *SAMLService) HandleSAMLAssertion(
ctx context.Context, ctx context.Context,
req *http.Request, req *http.Request,
) (*SAMLUserInfo, error) { ) (*SAMLUserInfo, error) {
relayStateToken := req.FormValue("RelayState")
if relayStateToken == "" {
return nil, fmt.Errorf("missing RelayState in SAML response")
}
var relayState coredata.SAMLRelayState
var samlRequest coredata.SAMLRequest
var config coredata.SAMLConfiguration
var org coredata.Organization
now := time.Now()
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)
}
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, err
}
samlResponseEncoded := req.FormValue("SAMLResponse") samlResponseEncoded := req.FormValue("SAMLResponse")
if samlResponseEncoded == "" { if samlResponseEncoded == "" {
return nil, fmt.Errorf("missing SAMLResponse in request") return nil, fmt.Errorf("missing SAMLResponse in request")
} }
scope := coredata.NewScope(org.TenantID) relayStateToken := req.FormValue("RelayState")
err = s.pg.WithConn( samlConfigIDParam := req.URL.Query().Get("c")
ctx, now := time.Now()
func(conn pg.Conn) error {
if err := config.LoadByID(ctx, conn, scope, relayState.SAMLConfigID); err != nil {
return fmt.Errorf("cannot load SAML configuration: %w", err)
}
return nil var config *coredata.SAMLConfiguration
}, var org *coredata.Organization
) var possibleRequestIDs []string
if err != nil { var err error
return nil, fmt.Errorf("cannot load SAML configuration: %w", err)
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{}
} }
if !config.Enabled { if !config.Enabled {
return nil, ErrSAMLDisabled{OrganizationID: config.OrganizationID} return nil, ErrSAMLDisabled{OrganizationID: config.OrganizationID}
} }
sp, err := s.GetServiceProvider(ctx, &config) sp, err := s.GetServiceProvider(ctx, config)
if err != nil { if err != nil {
return nil, ErrCannotCreateServiceProvider{Err: err} return nil, ErrCannotCreateServiceProvider{Err: err}
} }
@@ -502,7 +548,6 @@ func (s *SAMLService) HandleSAMLAssertion(
req.URL.Host = req.Host req.URL.Host = req.Host
} }
possibleRequestIDs := []string{samlRequest.ID}
assertion, err := sp.ParseResponse(req, possibleRequestIDs) assertion, err := sp.ParseResponse(req, possibleRequestIDs)
if err != nil { if err != nil {
return nil, fmt.Errorf( return nil, fmt.Errorf(
@@ -516,6 +561,7 @@ func (s *SAMLService) HandleSAMLAssertion(
if err := ValidateAssertion(assertion, s.GetEntityID(), now); err != nil { if err := ValidateAssertion(assertion, s.GetEntityID(), now); err != nil {
return nil, ErrCannotValidateAssertion{Err: err} return nil, ErrCannotValidateAssertion{Err: err}
} }
if assertion.ID != "" { if assertion.ID != "" {
var expiresAt time.Time var expiresAt time.Time
if assertion.Conditions != nil && !assertion.Conditions.NotOnOrAfter.IsZero() { if assertion.Conditions != nil && !assertion.Conditions.NotOnOrAfter.IsZero() {
@@ -525,22 +571,14 @@ func (s *SAMLService) HandleSAMLAssertion(
} }
scope := coredata.NewScope(org.TenantID) scope := coredata.NewScope(org.TenantID)
err = s.pg.WithTx( err = s.pg.WithTx(ctx, func(tx pg.Conn) error {
ctx, return PreventReplayAttack(ctx, tx, scope, assertion.ID, config.OrganizationID, expiresAt)
func(tx pg.Conn) error { })
if err := PreventReplayAttack(ctx, tx, scope, assertion.ID, relayState.OrganizationID, expiresAt); err != nil {
return fmt.Errorf("cannot prevent replay attack: %w", err)
}
return nil
},
)
if err != nil { if err != nil {
var replayAttackErr *coredata.ErrAssertionAlreadyUsed var replayAttackErr *coredata.ErrAssertionAlreadyUsed
if errors.As(err, &replayAttackErr) { if errors.As(err, &replayAttackErr) {
return nil, ErrReplayAttackDetected{AssertionID: assertion.ID, Err: replayAttackErr} return nil, ErrReplayAttackDetected{AssertionID: assertion.ID, Err: replayAttackErr}
} }
return nil, fmt.Errorf("cannot prevent replay attack: %w", err) return nil, fmt.Errorf("cannot prevent replay attack: %w", err)
} }
} }
@@ -576,8 +614,8 @@ func (s *SAMLService) HandleSAMLAssertion(
FullName: fullname, FullName: fullname,
Role: systemRole, Role: systemRole,
SAMLSubject: samlSubject, SAMLSubject: samlSubject,
OrganizationID: relayState.OrganizationID, OrganizationID: config.OrganizationID,
SAMLConfigID: relayState.SAMLConfigID, SAMLConfigID: config.ID,
}, nil }, nil
} }

View File

@@ -60,10 +60,13 @@ func SAMLACSHandler(samlSvc *authsvc.SAMLService, authSvc *authsvc.Service, auth
return return
} }
if r.FormValue("RelayState") == "" { relayState := r.FormValue("RelayState")
logger.WarnCtx(ctx, "missing RelayState") samlConfigID := r.URL.Query().Get("c")
http.Error(w, "missing RelayState", http.StatusBadRequest)
return if relayState != "" {
logger.InfoCtx(ctx, "processing SP-initiated SAML login", log.String("relay_state", relayState))
} else {
logger.InfoCtx(ctx, "processing IDP-initiated SAML login", log.String("config_id", samlConfigID))
} }
userInfo, err := samlSvc.HandleSAMLAssertion(ctx, r) userInfo, err := samlSvc.HandleSAMLAssertion(ctx, r)