mirror of
https://github.com/pocket-id/pocket-id.git
synced 2026-09-30 06:49:06 +02:00
fix: add 1min grace period to refresh token invalidation
This commit is contained in:
@@ -22,6 +22,7 @@ type OAuth2Session struct {
|
||||
Active bool
|
||||
RequestData string
|
||||
ExpiresAt *datatype.DateTime
|
||||
RotatedAt *datatype.DateTime
|
||||
}
|
||||
|
||||
func (OAuth2Session) TableName() string {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
+1
@@ -0,0 +1 @@
|
||||
ALTER TABLE oauth2_sessions DROP COLUMN rotated_at;
|
||||
+1
@@ -0,0 +1 @@
|
||||
ALTER TABLE oauth2_sessions ADD COLUMN rotated_at TIMESTAMPTZ;
|
||||
+1
@@ -0,0 +1 @@
|
||||
ALTER TABLE oauth2_sessions DROP COLUMN rotated_at;
|
||||
@@ -0,0 +1 @@
|
||||
ALTER TABLE oauth2_sessions ADD COLUMN rotated_at INTEGER;
|
||||
@@ -409,49 +409,6 @@ test('Refresh token fails when used for the wrong user', async ({ request }) =>
|
||||
expect(refreshResponse.status()).toBe(400);
|
||||
});
|
||||
|
||||
test('Using refresh token invalidates it for future use', async ({ request }) => {
|
||||
const { token, clientId, userId } = refreshTokens.filter((token) => !token.expired)[0];
|
||||
const clientSecret = 'w2mUeZISmEvIDMEDvpY0PnxQIpj1m3zY';
|
||||
|
||||
// Sign the refresh token
|
||||
const refreshToken = await request
|
||||
.post('/api/test/refreshtoken', {
|
||||
data: {
|
||||
rt: token,
|
||||
client: clientId,
|
||||
user: userId
|
||||
}
|
||||
})
|
||||
.then((r) => r.text());
|
||||
|
||||
// Perform the exchange
|
||||
await request.post('/api/oidc/token', {
|
||||
headers: {
|
||||
'Content-Type': 'application/x-www-form-urlencoded'
|
||||
},
|
||||
form: {
|
||||
grant_type: 'refresh_token',
|
||||
client_id: clientId,
|
||||
refresh_token: refreshToken,
|
||||
client_secret: clientSecret
|
||||
}
|
||||
});
|
||||
|
||||
// Try again
|
||||
const refreshResponse = await request.post('/api/oidc/token', {
|
||||
headers: {
|
||||
'Content-Type': 'application/x-www-form-urlencoded'
|
||||
},
|
||||
form: {
|
||||
grant_type: 'refresh_token',
|
||||
client_id: clientId,
|
||||
refresh_token: refreshToken,
|
||||
client_secret: clientSecret
|
||||
}
|
||||
});
|
||||
expect(refreshResponse.status()).toBe(400);
|
||||
});
|
||||
|
||||
test.describe('Introspection endpoint', () => {
|
||||
test('fails without client credentials', async ({ request }) => {
|
||||
const validAccessToken = await generateSeededOauthAccessToken(
|
||||
|
||||
Reference in New Issue
Block a user