fix: add 1min grace period to refresh token invalidation

This commit is contained in:
Elias Schneider
2026-09-29 22:56:21 +02:00
parent ca22253f48
commit b2d4faebf7
8 changed files with 215 additions and 54 deletions
+1
View File
@@ -22,6 +22,7 @@ type OAuth2Session struct {
Active bool
RequestData string
ExpiresAt *datatype.DateTime
RotatedAt *datatype.DateTime
}
func (OAuth2Session) TableName() string {
+173
View File
@@ -0,0 +1,173 @@
package oidc
import (
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"encoding/json"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"testing"
"time"
"github.com/gin-gonic/gin"
"github.com/ory/fosite"
"github.com/ory/fosite/compose"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
"github.com/pocket-id/pocket-id/backend/internal/model"
datatype "github.com/pocket-id/pocket-id/backend/internal/model/types"
testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing"
"github.com/pocket-id/pocket-id/backend/resources"
)
func TestRefreshTokenRotationGrace(t *testing.T) {
gin.SetMode(gin.TestMode)
const clientID, userID, requestID = "grace-client", "grace-user", "grace-request"
const issuer, secret = "https://issuer.example.com", "test-secret"
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
require.NoError(t, err)
globalSecret, err := DeriveGlobalSecret([]byte(secret))
require.NoError(t, err)
strategy := compose.NewOAuth2HMACStrategy(&fosite.Config{GlobalSecret: globalSecret})
for _, scenario := range []string{"retry", "concurrent", "reuse after grace", "expired token", "wrong client", "disabled user", "token revocation", "logout", "app revocation"} {
t.Run(scenario, func(t *testing.T) {
db := testutils.NewConcurrentDatabaseForTest(t)
store := NewStore(db, nil)
require.NoError(t, db.Create(&model.OidcClient{Base: model.Base{ID: clientID}, Name: "Grace client", IsPublic: true}).Error)
require.NoError(t, db.Create(&model.User{Base: model.Base{ID: userID}, Username: "grace"}).Error)
provider, err := newProvider(store, nil, testTokenSigner{key: key}, Config{BaseURL: issuer, TokenBaseURL: issuer, Secret: []byte(secret)}, nil)
require.NoError(t, err)
handler := newTokenHandler(provider, newClaimsService(db, nil, issuer, nil), nil)
// Seed an old grant so the grace period must start at rotation rather than token issuance
request := newTestRequester(requestID, clientID, userID, "grace-jti")
request.GetSession().SetExpiresAt(fosite.RefreshToken, time.Now().UTC().Add(time.Hour))
token, signature, err := strategy.GenerateRefreshToken(t.Context(), request)
require.NoError(t, err)
require.NoError(t, store.CreateRefreshTokenSession(t.Context(), signature, "original-access", request))
require.NoError(t, store.CreateAccessTokenSession(t.Context(), "original-access", request))
require.NoError(t, db.Model(&OAuth2Session{}).Where("kind = ? AND key = ?", sessionKindRefreshToken, signature).
Update("created_at", datatype.DateTime(time.Now().UTC().Add(-24*time.Hour))).Error)
// Exercise the real token endpoint, including authentication, user checks, rotation, and transactions
refresh := func(clientID, token string) *httptest.ResponseRecorder {
form := url.Values{"grant_type": {"refresh_token"}, "refresh_token": {token}, "client_id": {clientID}}
req := httptest.NewRequestWithContext(t.Context(), http.MethodPost, "/api/oidc/token", strings.NewReader(form.Encode()))
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = req
handler.token(c)
return rec
}
body := func(rec *httptest.ResponseRecorder, status int) map[string]any {
t.Helper()
require.Equal(t, status, rec.Code, rec.Body.String())
var result map[string]any
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &result))
return result
}
var first, second map[string]any
if scenario == "concurrent" {
start := make(chan struct{})
responses := make(chan *httptest.ResponseRecorder, 2)
for range 2 {
go func() {
<-start
responses <- refresh(clientID, token)
}()
}
close(start)
first = body(<-responses, http.StatusOK)
second = body(<-responses, http.StatusOK)
} else {
first = body(refresh(clientID, token), http.StatusOK)
}
if scenario == "retry" {
second = body(refresh(clientID, token), http.StatusOK)
}
rotated, err := store.getSession(t.Context(), sessionKindRefreshToken, signature)
require.NoError(t, err)
require.False(t, rotated.Active)
require.NotNil(t, rotated.RotatedAt)
_, err = store.GetAccessTokenSession(t.Context(), "original-access", nil)
require.ErrorIs(t, err, fosite.ErrNotFound)
switch scenario {
case "retry", "concurrent":
require.NotEqual(t, first["refresh_token"], second["refresh_token"])
afterRetry, err := store.getSession(t.Context(), sessionKindRefreshToken, signature)
require.NoError(t, err)
require.Equal(t, rotated.RotatedAt, afterRetry.RotatedAt)
for _, issued := range []map[string]any{first, second} {
_, err := store.GetAccessTokenSession(t.Context(), provider.accessToken.AccessTokenSignature(t.Context(), issued["access_token"].(string)), nil)
require.NoError(t, err)
body(refresh(clientID, issued["refresh_token"].(string)), http.StatusOK)
}
case "reuse after grace":
second = body(refresh(clientID, token), http.StatusOK)
require.NoError(t, db.Model(&OAuth2Session{}).Where("kind = ? AND key = ?", sessionKindRefreshToken, signature).
Update("rotated_at", datatype.DateTime(time.Now().UTC().Add(-refreshTokenGracePeriod-time.Second))).Error)
require.ErrorIs(t, store.RotateRefreshToken(t.Context(), requestID, signature), fosite.ErrInactiveToken)
require.Equal(t, "invalid_grant", body(refresh(clientID, token), http.StatusBadRequest)["error"])
for _, issued := range []map[string]any{first, second} {
_, err := store.GetAccessTokenSession(t.Context(), provider.accessToken.AccessTokenSignature(t.Context(), issued["access_token"].(string)), nil)
require.ErrorIs(t, err, fosite.ErrNotFound)
require.Equal(t, "invalid_grant", body(refresh(clientID, issued["refresh_token"].(string)), http.StatusBadRequest)["error"])
}
case "expired token":
request.GetSession().SetExpiresAt(fosite.RefreshToken, time.Now().UTC().Add(-time.Hour))
data, err := store.encodeRequester(request)
require.NoError(t, err)
require.NoError(t, db.Model(&OAuth2Session{}).Where("kind = ? AND key = ?", sessionKindRefreshToken, signature).Update("request_data", data).Error)
require.Equal(t, "invalid_grant", body(refresh(clientID, token), http.StatusBadRequest)["error"])
body(refresh(clientID, first["refresh_token"].(string)), http.StatusOK)
case "wrong client":
require.NoError(t, db.Create(&model.OidcClient{Base: model.Base{ID: "other-client"}, Name: "Other", IsPublic: true}).Error)
require.Equal(t, "invalid_grant", body(refresh("other-client", token), http.StatusBadRequest)["error"])
body(refresh(clientID, first["refresh_token"].(string)), http.StatusOK)
case "disabled user":
require.NoError(t, db.Model(&model.User{}).Where("id = ?", userID).Update("disabled", true).Error)
require.Equal(t, "invalid_grant", body(refresh(clientID, token), http.StatusBadRequest)["error"])
default:
// Revoking the grant must cancel grace for ancestors as well as their descendants
switch scenario {
case "token revocation":
require.NoError(t, store.RevokeRefreshToken(t.Context(), requestID))
case "logout":
require.NoError(t, store.RevokeSessionsByIDTokenHint(t.Context(), userID, clientID, "grace-jti"))
case "app revocation":
require.NoError(t, RevokeUserClientSessions(t.Context(), db, userID, clientID))
}
require.ErrorIs(t, store.RotateRefreshToken(t.Context(), requestID, signature), fosite.ErrInactiveToken)
require.Equal(t, "invalid_grant", body(refresh(clientID, token), http.StatusBadRequest)["error"])
require.Equal(t, "invalid_grant", body(refresh(clientID, first["refresh_token"].(string)), http.StatusBadRequest)["error"])
}
})
}
}
func TestRefreshTokenGraceMigrationPreservesExistingSessions(t *testing.T) {
const previousVersion = 20260923183637
db := testutils.NewDatabaseForTestWithMigrationSeed(t, previousVersion, func(t *testing.T, db *gorm.DB) {
require.NoError(t, db.Create(&model.OidcClient{Base: model.Base{ID: "old-client"}, Name: "Old client"}).Error)
require.NoError(t, db.Exec(`INSERT INTO oauth2_sessions (id, created_at, kind, key, request_id, client_id, active, request_data) VALUES ('old-refresh', 1, 'refresh_token', 'old-signature', 'old-request', 'old-client', false, '{}')`).Error)
})
var session OAuth2Session
require.NoError(t, db.First(&session, "id = ?", "old-refresh").Error)
require.False(t, session.Active)
require.Nil(t, session.RotatedAt)
down, err := resources.FS.ReadFile("migrations/sqlite/20260929120000_refresh_token_rotation_grace.down.sql")
require.NoError(t, err)
require.NoError(t, db.Exec(string(down)).Error)
var count int64
require.NoError(t, db.Table("oauth2_sessions").Where("id = ?", "old-refresh").Count(&count).Error)
require.EqualValues(t, 1, count)
}
+37 -11
View File
@@ -31,6 +31,7 @@ const (
sessionKindPAR = "par"
sessionKindDeviceCode = "device_code"
sessionKindUserCode = "user_code"
refreshTokenGracePeriod = time.Minute
)
var (
@@ -364,11 +365,16 @@ func (s *Store) CreateRefreshTokenSession(ctx context.Context, signature string,
}
func (s *Store) GetRefreshTokenSession(ctx context.Context, signature string, _ fosite.Session) (fosite.Requester, error) {
request, active, err := s.getRequesterSession(ctx, sessionKindRefreshToken, signature)
session, err := s.getSession(ctx, sessionKindRefreshToken, signature)
if err != nil {
return nil, err
}
if !active {
request, err := s.decodeRequester(ctx, session.RequestData)
if err != nil {
return nil, err
}
// Allow retries briefly after rotation while keeping explicitly revoked tokens inactive
if !session.Active && (session.RotatedAt == nil || !time.Now().UTC().Before(session.RotatedAt.ToTime().Add(refreshTokenGracePeriod))) {
return request, fosite.ErrInactiveToken
}
return request, nil
@@ -379,10 +385,29 @@ func (s *Store) DeleteRefreshTokenSession(ctx context.Context, signature string)
}
func (s *Store) RotateRefreshToken(ctx context.Context, requestID string, refreshTokenSignature string) error {
if err := s.deactivateSession(ctx, sessionKindRefreshToken, refreshTokenSignature); err != nil {
// Atomically start the grace period once so parallel refreshes cannot extend it or revive revoked tokens
now := time.Now().UTC()
result := s.dbFor(ctx).
Model(&OAuth2Session{}).
Where("kind = ? AND key = ? AND request_id = ?", sessionKindRefreshToken, refreshTokenSignature, requestID).
Where("active = ? OR rotated_at > ?", true, datatype.DateTime(now.Add(-refreshTokenGracePeriod))).
Updates(map[string]any{
"active": false,
"rotated_at": gorm.Expr("COALESCE(rotated_at, ?)", datatype.DateTime(now)),
})
if result.Error != nil {
return result.Error
}
if result.RowsAffected == 0 {
return fosite.ErrInactiveToken
}
// Revoke only the access token paired with this refresh token so concurrent refreshes keep their access tokens
session, err := s.getSession(ctx, sessionKindRefreshToken, refreshTokenSignature)
if err != nil {
return err
}
return s.RevokeAccessToken(ctx, requestID)
return s.DeleteAccessTokenSession(ctx, session.AccessTokenSignature)
}
// Satisfies fositeoauth2.TokenRevocationStorage
@@ -391,7 +416,7 @@ func (s *Store) RevokeRefreshToken(ctx context.Context, requestID string) error
return s.dbFor(ctx).
Model(&OAuth2Session{}).
Where("kind = ? AND request_id = ?", sessionKindRefreshToken, requestID).
Update("active", false).
Updates(map[string]any{"active": false, "rotated_at": nil}).
Error
}
@@ -403,7 +428,7 @@ func (s *Store) RevokeAccessToken(ctx context.Context, requestID string) error {
}
func (s *Store) RevokeSessionsByIDTokenHint(ctx context.Context, userID, clientID, idTokenJTI string) error {
_, jtiMatches, err := s.findActiveRefreshTokenRequestIDsForUserClient(ctx, userID, clientID, idTokenJTI)
_, jtiMatches, err := s.findRefreshTokenRequestIDsForUserClient(ctx, userID, clientID, idTokenJTI)
if err != nil {
return err
}
@@ -413,19 +438,20 @@ func (s *Store) RevokeSessionsByIDTokenHint(ctx context.Context, userID, clientI
func RevokeUserClientSessions(ctx context.Context, db *gorm.DB, userID, clientID string) error {
s := NewStore(db, nil)
requestIDs, _, err := s.findActiveRefreshTokenRequestIDsForUserClient(ctx, userID, clientID, "")
requestIDs, _, err := s.findRefreshTokenRequestIDsForUserClient(ctx, userID, clientID, "")
if err != nil {
return err
}
return s.revokeRequestIDs(ctx, requestIDs)
}
// findActiveRefreshTokenRequestIDsForUserClient returns request IDs for active refresh-token sessions belonging to the user and client, plus the subset matching the optional ID token JTI
func (s *Store) findActiveRefreshTokenRequestIDsForUserClient(ctx context.Context, userID, clientID, idTokenJTI string) (candidates []string, jtiMatches []string, err error) {
// findRefreshTokenRequestIDsForUserClient returns request IDs for usable refresh-token sessions belonging to the user and client, plus the subset matching the optional ID token JTI
func (s *Store) findRefreshTokenRequestIDsForUserClient(ctx context.Context, userID, clientID, idTokenJTI string) (candidates []string, jtiMatches []string, err error) {
var sessions []OAuth2Session
query := s.dbFor(ctx).
Select("request_id", "request_data").
Where("kind = ? AND active = ? AND client_id = ?", sessionKindRefreshToken, true, clientID)
Where("kind = ? AND client_id = ?", sessionKindRefreshToken, clientID).
Where("active = ? OR rotated_at > ?", true, datatype.DateTime(time.Now().UTC().Add(-refreshTokenGracePeriod)))
// Filter by the user ID stored in the JSON request data
switch query.Name() {
@@ -472,7 +498,7 @@ func (s *Store) revokeRequestIDs(ctx context.Context, requestIDs []string) error
if err := s.dbFor(ctx).
Model(&OAuth2Session{}).
Where("kind = ? AND request_id IN ?", sessionKindRefreshToken, requestIDs).
Update("active", false).
Updates(map[string]any{"active": false, "rotated_at": nil}).
Error; err != nil {
return err
}