Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
57 changes: 50 additions & 7 deletions backend/internal/auth/token.go
Original file line number Diff line number Diff line change
Expand Up @@ -425,6 +425,14 @@ func (tm *TokenManager) RefreshTokenPair(ctx context.Context, refreshTokenValue
// only a revocation (or an explicit logout) takes it away. If it is gone, the
// lineage was declared compromised — drop what we issued and refuse, the same
// answer the losing request got.
//
// Seeing the witness is not enough on its own (#725). On Postgres a revoking
// DELETE reads from the snapshot taken when it starts and its deletions stay
// invisible until it commits, so we can store our token after that snapshot
// and still find the witness here. The revokers close that case by sweeping
// until a pass finds nothing (sweepRefreshTokens): our token was committed
// before this read, this read came before their first pass committed, so
// their next pass sees our token and takes it.
if !tm.tokenExists(ctx, refreshToken.ID) {
tm.revokeFamily(ctx, refreshToken.FamilyID)
return nil, ErrRefreshTokenReuse
Expand Down Expand Up @@ -532,7 +540,44 @@ func (tm *TokenManager) revokeFamily(ctx context.Context, familyID uuid.UUID) {
if familyID == uuid.Nil {
return
}
tm.db.WithContext(ctx).Where("family_id = ?", familyID).Delete(&RefreshToken{})
_ = tm.sweepRefreshTokens(ctx, "family_id = ?", familyID)
}

// maxRevocationSweeps bounds sweepRefreshTokens. Each extra pass is only needed
// when a rotation stored a token behind the previous one, which takes a client
// round trip per step, so a handful is far more than a real race produces.
const maxRevocationSweeps = 5

// sweepRefreshTokens deletes the refresh tokens matching the condition, and
// repeats until a pass removes nothing (#725).
//
// One DELETE is not enough on Postgres. Under READ COMMITTED it only sees rows
// committed before it started, so a rotation that stores its successor while
// the DELETE runs leaves that successor behind. The rotation cannot catch this
// itself: it checks its witness row after storing the successor, and it still
// sees the witness because the DELETE has not committed yet.
//
// The ordering argument. Let P be the last pass, the one that removed nothing.
// When P started, every row revoked here was gone, witnesses included.
// - A successor committed before P started is visible to P, so it was
// already deleted.
// - A successor committed after P started is checked against its witness
// after that commit, so after the earlier passes committed. The rotation
// finds no witness, revokes the family itself and refuses (RefreshTokenPair).
//
// A pass that removes nothing costs one indexed DELETE, on revocation only. The
// refresh path takes no lock and runs no extra query.
func (tm *TokenManager) sweepRefreshTokens(ctx context.Context, query string, args ...interface{}) error {
for i := 0; i < maxRevocationSweeps; i++ {
res := tm.db.WithContext(ctx).Where(query, args...).Delete(&RefreshToken{})
if res.Error != nil {
return res.Error
}
if res.RowsAffected == 0 {
return nil
}
}
return nil
}

// PruneExpiredTokens removes refresh tokens past their TTL (both live and spent).
Expand All @@ -559,9 +604,8 @@ func (tm *TokenManager) RevokeRefreshToken(ctx context.Context, refreshTokenValu

// RevokeAllUserTokens revokes all refresh tokens for a user
func (tm *TokenManager) RevokeAllUserTokens(ctx context.Context, userID uuid.UUID) error {
result := tm.db.WithContext(ctx).Where("user_id = ?", userID).Delete(&RefreshToken{})
if result.Error != nil {
return fmt.Errorf("failed to revoke user tokens: %w", result.Error)
if err := tm.sweepRefreshTokens(ctx, "user_id = ?", userID); err != nil {
return fmt.Errorf("failed to revoke user tokens: %w", err)
}
return nil
}
Expand All @@ -572,9 +616,8 @@ func (tm *TokenManager) RevokeAllUserTokens(ctx context.Context, userID uuid.UUI
// organization, never their sessions elsewhere (#831). Account-level events
// (password change or reset) use RevokeAllUserTokens instead.
func (tm *TokenManager) RevokeUserTokensInTenant(ctx context.Context, userID, tenantID uuid.UUID) error {
result := tm.db.WithContext(ctx).Where("user_id = ? AND tenant_id = ?", userID, tenantID).Delete(&RefreshToken{})
if result.Error != nil {
return fmt.Errorf("failed to revoke user tokens in tenant: %w", result.Error)
if err := tm.sweepRefreshTokens(ctx, "user_id = ? AND tenant_id = ?", userID, tenantID); err != nil {
return fmt.Errorf("failed to revoke user tokens in tenant: %w", err)
}
return nil
}
Expand Down
153 changes: 153 additions & 0 deletions backend/internal/auth/token_pg_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,153 @@
// Copyright (c) 2026 OpenDefender Contributors
// SPDX-License-Identifier: AGPL-3.0-only
// This program is free software: you can redistribute it and/or modify it under
// the terms of the GNU Affero General Public License v3.0 (see LICENSE).

package auth

import (
"context"
"crypto/rand"
"crypto/rsa"
"fmt"
"net/url"
"os"
"strings"
"testing"
"time"

"github.com/google/uuid"
"github.com/stretchr/testify/require"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/logger"

authpkg "github.com/opendefender/openrisk/pkg/auth"
)

// newPgTokenHarness opens DATABASE_URL on a schema of its own, so the test's
// refresh_tokens table never meets the real one, and drops it afterwards.
func newPgTokenHarness(t *testing.T) (*TokenManager, *gorm.DB) {
t.Helper()
dsn := os.Getenv("DATABASE_URL")
if dsn == "" {
t.Skip("DATABASE_URL not set")
}
admin, err := gorm.Open(postgres.Open(dsn), &gorm.Config{Logger: logger.Discard})
require.NoError(t, err)
schema := "auth_test_" + strings.ReplaceAll(uuid.NewString(), "-", "")[:12]
require.NoError(t, admin.Exec("CREATE SCHEMA "+schema).Error)
t.Cleanup(func() {
admin.Exec("DROP SCHEMA " + schema + " CASCADE")
if sqlDB, err := admin.DB(); err == nil {
sqlDB.Close()
}
})

u, err := url.Parse(dsn)
require.NoError(t, err)
q := u.Query()
q.Set("search_path", schema)
u.RawQuery = q.Encode()
db, err := gorm.Open(postgres.Open(u.String()), &gorm.Config{Logger: logger.Discard})
require.NoError(t, err)
t.Cleanup(func() {
if sqlDB, err := db.DB(); err == nil {
sqlDB.Close()
}
})
require.NoError(t, db.AutoMigrate(&RefreshToken{}))

priv, err := rsa.GenerateKey(rand.Reader, 2048)
require.NoError(t, err)
return NewTokenManager(db, &authpkg.RSAKeys{PrivateKey: priv, PublicKey: &priv.PublicKey}), db
}

// waitForLockWaiter blocks until some backend is waiting on a row lock in the
// test's table: the revoking DELETE has started, taken its snapshot, and is
// parked behind the lock the test holds.
func waitForLockWaiter(t *testing.T, db *gorm.DB) {
t.Helper()
deadline := time.Now().Add(10 * time.Second)
for time.Now().Before(deadline) {
var n int64
require.NoError(t, db.Raw(`SELECT count(*) FROM pg_stat_activity
WHERE wait_event_type = 'Lock' AND query ILIKE 'DELETE FROM "refresh_tokens"%'`).Scan(&n).Error)
if n > 0 {
return
}
time.Sleep(10 * time.Millisecond)
}
t.Fatal("the revoking DELETE never blocked on the held lock")
}

// TestRevocation_RacingRotation_Postgres drives, on Postgres, the interleaving
// SQLite cannot produce (#725). Under READ COMMITTED a DELETE reads from the
// snapshot taken when it starts, and its deletions stay invisible until it
// commits. So:
//
// 1. the rotation claims the presented token (the witness);
// 2. a revocation starts its DELETE, and a lock held by the test parks it
// before it commits;
// 3. the rotation stores its successor and still sees the witness, because the
// DELETE has not committed, so it hands the successor out;
// 4. the lock is released, and the DELETE commits without the successor, which
// was not in its snapshot.
//
// Every revoker must leave no token of what it revoked, and the token handed out
// in step 3 must be refused.
func TestRevocation_RacingRotation_Postgres(t *testing.T) {
cases := []struct {
name string
revoke func(ctx context.Context, tm *TokenManager, rt RefreshToken) error
}{
{"family", func(ctx context.Context, tm *TokenManager, rt RefreshToken) error {
tm.revokeFamily(ctx, rt.FamilyID)
return nil
}},
{"all user tokens", func(ctx context.Context, tm *TokenManager, rt RefreshToken) error {
return tm.RevokeAllUserTokens(ctx, rt.UserID)
}},
{"user tokens in tenant", func(ctx context.Context, tm *TokenManager, rt RefreshToken) error {
return tm.RevokeUserTokensInTenant(ctx, rt.UserID, rt.TenantID)
}},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
tm, db := newPgTokenHarness(t)
ctx := context.Background()
userID, orgID := uuid.New(), uuid.New()

pair, err := tm.GenerateTokenPair(ctx, userID, orgID, nil, []string{"*"}, nil, DeviceContext{})
require.NoError(t, err)
var witness RefreshToken
require.NoError(t, db.First(&witness).Error)

lock := db.Begin()
defer lock.Rollback()
revoked := make(chan error, 1)

// The org resolver runs between the claim and the insert.
tm.SetOrgSessionResolver(func(_ context.Context, _ uuid.UUID, org uuid.UUID) (*SessionClaims, error) {
var held RefreshToken
if err := lock.Raw("SELECT * FROM refresh_tokens WHERE id = ? FOR UPDATE", witness.ID).Scan(&held).Error; err != nil {
return nil, fmt.Errorf("lock witness: %w", err)
}
go func() { revoked <- tc.revoke(ctx, tm, witness) }()
waitForLockWaiter(t, db)
return &SessionClaims{TenantID: org, Permissions: []string{"*"}}, nil
})

issued, rotErr := tm.RefreshTokenPair(ctx, pair.RefreshToken, DeviceContext{})

require.NoError(t, lock.Commit().Error)
require.NoError(t, <-revoked)

require.Equal(t, int64(0), countTokens(t, db), "no token may outlive the revocation it raced")
if rotErr == nil {
_, err := tm.RefreshTokenPair(ctx, issued.RefreshToken, DeviceContext{})
require.ErrorIs(t, err, ErrRefreshTokenInvalid, "the token handed out mid-revocation must be dead")
}
})
}
}
29 changes: 29 additions & 0 deletions backend/internal/auth/token_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -417,3 +417,32 @@ func TestSuccessorSecret_IsKeyedAndBound(t *testing.T) {
require.NotEqual(t, presented, a.successorSecret(presented, family))
require.Len(t, a.successorSecret(presented, family), 64, "same shape as a random token")
}

// TestSweepRefreshTokens_TakesARowStoredBehindIt models what a Postgres DELETE
// misses (#725): a row committed after the pass has started. The sweep must
// take it on its next pass, and stop at the first pass that removes nothing.
func TestSweepRefreshTokens_TakesARowStoredBehindIt(t *testing.T) {
tm, db, _ := newTokenHarness(t)
ctx := context.Background()
userID, orgID := uuid.New(), uuid.New()

_, err := tm.GenerateTokenPair(ctx, userID, orgID, nil, []string{"*"}, nil, DeviceContext{})
require.NoError(t, err)
var family uuid.UUID
require.NoError(t, db.Model(&RefreshToken{}).Select("family_id").Row().Scan(&family))

passes := 0
require.NoError(t, db.Callback().Delete().After("gorm:delete").Register("test:late_row", func(tx *gorm.DB) {
passes++
if passes == 1 {
// A rotation's successor, landing just behind the first pass.
late := RefreshToken{UserID: userID, TenantID: orgID, FamilyID: family,
TokenHash: hashToken(uuid.NewString()), ExpiresAt: time.Now().Add(time.Hour)}
require.NoError(t, db.Session(&gorm.Session{NewDB: true}).Create(&late).Error)
}
}))

tm.revokeFamily(ctx, family)
require.Equal(t, int64(0), countTokens(t, db), "the row stored behind the first pass must be swept")
require.Equal(t, 3, passes, "two passes that removed a row, then one that removed nothing")
}
Loading
Loading