From b92d78e4cf0af2c06a909dc9b03a40b4e1483269 Mon Sep 17 00:00:00 2001 From: Bryan Frimin Date: Wed, 29 Oct 2025 20:36:52 +0100 Subject: [PATCH] Remove race condition on relay attack Signed-off-by: Bryan Frimin --- pkg/auth/saml_validator.go | 18 +++++++----------- pkg/coredata/saml_assertion.go | 28 ---------------------------- 2 files changed, 7 insertions(+), 39 deletions(-) diff --git a/pkg/auth/saml_validator.go b/pkg/auth/saml_validator.go index b3267bdd1..e58af7a16 100644 --- a/pkg/auth/saml_validator.go +++ b/pkg/auth/saml_validator.go @@ -16,12 +16,14 @@ package auth import ( "context" + "errors" "fmt" "time" "github.com/crewjam/saml" "github.com/getprobo/probo/pkg/coredata" "github.com/getprobo/probo/pkg/gid" + "github.com/jackc/pgx/v5/pgconn" "go.gearno.de/kit/pg" ) @@ -33,18 +35,8 @@ func PreventReplayAttack( organizationID gid.GID, expiresAt time.Time, ) error { - var assertion coredata.SAMLAssertion - exists, err := assertion.CheckExists(ctx, conn, assertionID) - if err != nil { - return fmt.Errorf("cannot check assertion ID: %w", err) - } - - if exists { - return coredata.ErrAssertionAlreadyUsed{AssertionID: assertionID} - } - now := time.Now() - assertion = coredata.SAMLAssertion{ + assertion := coredata.SAMLAssertion{ ID: assertionID, OrganizationID: organizationID, UsedAt: now, @@ -52,6 +44,10 @@ func PreventReplayAttack( } if err := assertion.Insert(ctx, conn, scope); err != nil { + var pgErr *pgconn.PgError + if errors.As(err, &pgErr) && pgErr.Code == "23505" && pgErr.ConstraintName == "auth_saml_assertions_pkey" { + return coredata.ErrAssertionAlreadyUsed{AssertionID: assertionID} + } return fmt.Errorf("cannot store assertion ID: %w", err) } diff --git a/pkg/coredata/saml_assertion.go b/pkg/coredata/saml_assertion.go index 7d9f645b8..d687bf23b 100644 --- a/pkg/coredata/saml_assertion.go +++ b/pkg/coredata/saml_assertion.go @@ -39,34 +39,6 @@ func (e ErrAssertionAlreadyUsed) Error() string { return fmt.Sprintf("assertion ID %q has already been used (replay attack)", e.AssertionID) } -func (s *SAMLAssertion) CheckExists( - ctx context.Context, - conn pg.Conn, - assertionID string, -) (bool, error) { - query := ` -SELECT id -FROM auth_saml_assertions -WHERE id = @id -LIMIT 1 -` - - rows, err := conn.Query(ctx, query, pgx.NamedArgs{"id": assertionID}) - if err != nil { - return false, fmt.Errorf("cannot query saml_assertions: %w", err) - } - - _, err = pgx.CollectOneRow(rows, pgx.RowTo[string]) - if err == nil { - return true, nil - } - if err == pgx.ErrNoRows { - return false, nil - } - - return false, fmt.Errorf("cannot collect saml_assertion: %w", err) -} - func (s *SAMLAssertion) Insert( ctx context.Context, conn pg.Conn,