package auth import ( "context" "fmt" "sync" "testing" "github.com/alicebob/miniredis/v2" "github.com/gochat/gochat/internal/config" "github.com/redis/go-redis/v9" "github.com/stretchr/testify/require" ) func TestRefreshTokenStoreScopesTokensByClient(t *testing.T) { store := NewRefreshTokenStore(nil, &config.JWTConfig{RefreshExpiryHours: 24}) ctx := context.Background() require.NoError(t, store.StoreForClient(ctx, 7, "chrome", "chrome-token")) require.NoError(t, store.StoreForClient(ctx, 7, "firefox", "firefox-token")) valid, err := store.ValidateForClient(ctx, 7, "chrome", "chrome-token") require.NoError(t, err) require.True(t, valid) require.NoError(t, store.RevokeClient(ctx, 7, "chrome")) valid, err = store.ValidateForClient(ctx, 7, "chrome", "chrome-token") require.NoError(t, err) require.False(t, valid) valid, err = store.ValidateForClient(ctx, 7, "firefox", "firefox-token") require.NoError(t, err) require.True(t, valid) } func TestRefreshTokenStoreRevokeUserRedisDoesNotMatchSimilarUserIDs(t *testing.T) { mr := miniredis.RunT(t) store := NewRefreshTokenStore(redis.NewClient(&redis.Options{Addr: mr.Addr()}), &config.JWTConfig{RefreshExpiryHours: 24}) ctx := context.Background() require.NoError(t, store.StoreForClient(ctx, 7, "browser", "user-7")) require.NoError(t, store.StoreForClient(ctx, 70, "browser", "user-70")) require.NoError(t, store.RevokeUser(ctx, 7)) valid, err := store.ValidateForClient(ctx, 70, "browser", "user-70") require.NoError(t, err) require.True(t, valid) } func TestRefreshTokenStoreRevokeUserRemovesAllClients(t *testing.T) { store := NewRefreshTokenStore(nil, &config.JWTConfig{RefreshExpiryHours: 24}) ctx := context.Background() require.NoError(t, store.Store(ctx, 7, "legacy")) require.NoError(t, store.StoreForClient(ctx, 7, "chrome", "chrome-token")) require.NoError(t, store.StoreForClient(ctx, 7, "mobile", "mobile-token")) require.NoError(t, store.StoreForClient(ctx, 8, "chrome", "other-user")) require.NoError(t, store.RevokeUser(ctx, 7)) for clientID, token := range map[string]string{"": "legacy", "chrome": "chrome-token", "mobile": "mobile-token"} { valid, err := store.ValidateForClient(ctx, 7, clientID, token) require.NoError(t, err) require.False(t, valid) } valid, err := store.ValidateForClient(ctx, 8, "chrome", "other-user") require.NoError(t, err) require.True(t, valid) } func TestRefreshTokenStoreConcurrentDoubleRefreshAllowsOneWinner(t *testing.T) { for _, backend := range []string{"memory", "redis"} { t.Run(backend, func(t *testing.T) { var rdb *redis.Client if backend == "redis" { mr := miniredis.RunT(t) rdb = redis.NewClient(&redis.Options{Addr: mr.Addr()}) t.Cleanup(func() { require.NoError(t, rdb.Close()) }) } store := NewRefreshTokenStore(rdb, &config.JWTConfig{RefreshExpiryHours: 24}) ctx := context.Background() require.NoError(t, store.StoreForClient(ctx, 7, "browser", "old-token")) start := make(chan struct{}) results := make([]bool, 2) errs := make([]error, 2) var wg sync.WaitGroup for i := range results { wg.Add(1) go func(i int) { defer wg.Done() <-start results[i], errs[i] = store.CompareAndSwapForClient(ctx, 7, "browser", "old-token", fmt.Sprintf("new-token-%d", i)) }(i) } close(start) wg.Wait() require.NoError(t, errs[0]) require.NoError(t, errs[1]) require.NotEqual(t, results[0], results[1]) }) } }