Skip to content

Commit ca2cb7a

Browse files
committed
fix: prevent stale subscription cache updates
1 parent f010d6e commit ca2cb7a

4 files changed

Lines changed: 166 additions & 15 deletions

File tree

src/hooks/use-subscription-status.ts

Lines changed: 18 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -67,7 +67,7 @@ export const hasActiveSubscription = (
6767
*/
6868
export function useSubscriptionStatus() {
6969
const { user, isLoaded } = useUser()
70-
const [expirationTick, setExpirationTick] = useState(0)
70+
const [, setExpirationTick] = useState(0)
7171

7272
const publicMetadata = (user?.publicMetadata ?? {}) as Record<string, unknown>
7373
const rawChatStatus = publicMetadata['chat_subscription_status']
@@ -79,14 +79,23 @@ export function useSubscriptionStatus() {
7979

8080
useEffect(() => {
8181
if (expirationTime === null) return
82-
const delay = expirationTime - Date.now()
83-
if (delay <= 0) return
84-
const timeout = window.setTimeout(
85-
() => setExpirationTick((tick) => tick + 1),
86-
Math.min(delay + 1, MAX_TIMEOUT_MS),
87-
)
88-
return () => window.clearTimeout(timeout)
89-
}, [expirationTime, expirationTick])
82+
let timeout: number | undefined
83+
const scheduleExpiration = () => {
84+
const delay = expirationTime - Date.now()
85+
if (delay <= 0) {
86+
setExpirationTick((tick) => tick + 1)
87+
return
88+
}
89+
timeout = window.setTimeout(
90+
scheduleExpiration,
91+
Math.min(delay + 1, MAX_TIMEOUT_MS),
92+
)
93+
}
94+
scheduleExpiration()
95+
return () => {
96+
if (timeout !== undefined) window.clearTimeout(timeout)
97+
}
98+
}, [expirationTime])
9099

91100
const chatSubscriptionActive =
92101
isLoaded &&

src/services/inference/tinfoil-client.ts

Lines changed: 35 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,7 @@ let cachedSessionTokenWasAuthenticated = false
3838
let cachedRateLimit: RateLimitInfo | null = null
3939
let remainingBeforeRequest: number | null = null
4040
let refreshInFlight: Promise<void> | null = null
41+
let sessionCacheGeneration = 0
4142

4243
function dispatchRateLimitUpdate(): void {
4344
if (typeof window !== 'undefined') {
@@ -90,6 +91,7 @@ function surfaceHourlyLimit(parsedError: ServerErrorBody | null): never {
9091
// through the opaque path.
9192
async function fetchChatJWT(
9293
authBearer: string,
94+
cacheGeneration: number,
9395
): Promise<{ key: string; expiresAt: number | null } | null> {
9496
let response: Response
9597
try {
@@ -124,6 +126,7 @@ async function fetchChatJWT(
124126

125127
const parsedError = parseErrorBody(await response.text())
126128
if (isHourlyLimit(response.status, parsedError)) {
129+
if (cacheGeneration !== sessionCacheGeneration) return null
127130
surfaceHourlyLimit(parsedError)
128131
}
129132
return null
@@ -134,6 +137,8 @@ async function fetchSessionToken(): Promise<string> {
134137
return DEV_API_KEY
135138
}
136139

140+
const cacheGeneration = sessionCacheGeneration
141+
137142
// If the user was previously signed in, wait for Clerk to initialize
138143
// the auth token manager before fetching — otherwise we'd get an
139144
// anonymous free-tier key that gets cached until expiry.
@@ -164,6 +169,9 @@ async function fetchSessionToken(): Promise<string> {
164169
)
165170
}
166171
}
172+
if (cacheGeneration !== sessionCacheGeneration) {
173+
return fetchSessionToken()
174+
}
167175
const usedAuthHeader = authBearer !== null
168176

169177
// If the cached token was fetched anonymously but we now have an
@@ -196,7 +204,10 @@ async function fetchSessionToken(): Promise<string> {
196204
// Anonymous users (and signed-in users without an active subscription) fall
197205
// back to the opaque /api/keys/chat path below.
198206
if (authBearer) {
199-
const jwt = await fetchChatJWT(authBearer)
207+
const jwt = await fetchChatJWT(authBearer, cacheGeneration)
208+
if (cacheGeneration !== sessionCacheGeneration) {
209+
return fetchSessionToken()
210+
}
200211
if (jwt !== null) {
201212
cachedSessionToken = jwt.key
202213
cachedSessionTokenWasAuthenticated = true
@@ -216,9 +227,15 @@ async function fetchSessionToken(): Promise<string> {
216227
const response = await fetch(`${API_BASE_URL}/api/keys/chat`, {
217228
headers,
218229
})
230+
if (cacheGeneration !== sessionCacheGeneration) {
231+
return fetchSessionToken()
232+
}
219233

220234
if (!response.ok) {
221235
const errorText = await response.text()
236+
if (cacheGeneration !== sessionCacheGeneration) {
237+
return fetchSessionToken()
238+
}
222239
logError('Failed to fetch session token from server', undefined, {
223240
component: 'tinfoil-client',
224241
action: 'fetchSessionToken',
@@ -242,6 +259,9 @@ async function fetchSessionToken(): Promise<string> {
242259
}
243260

244261
const data = await response.json()
262+
if (cacheGeneration !== sessionCacheGeneration) {
263+
return fetchSessionToken()
264+
}
245265
cachedSessionToken = data.key
246266
cachedSessionTokenWasAuthenticated = usedAuthHeader
247267
if (data.expires_at) {
@@ -295,14 +315,16 @@ export function snapshotAndDecrementRemaining(): void {
295315
export async function refreshRateLimit(): Promise<void> {
296316
if (refreshInFlight) return refreshInFlight
297317

298-
refreshInFlight = (async () => {
318+
const refresh = (async () => {
319+
const refreshGeneration = sessionCacheGeneration
299320
const snapshot = remainingBeforeRequest
300321
remainingBeforeRequest = null
301322
cachedSessionToken = null
302323
cachedSessionTokenExpiresAt = null
303324
try {
304325
await fetchSessionToken()
305326
if (
327+
refreshGeneration === sessionCacheGeneration &&
306328
snapshot !== null &&
307329
cachedRateLimit &&
308330
cachedRateLimit.remaining >= snapshot
@@ -318,31 +340,38 @@ export async function refreshRateLimit(): Promise<void> {
318340
component: 'tinfoil-client',
319341
action: 'refreshRateLimit',
320342
})
321-
} finally {
322-
refreshInFlight = null
323343
}
324344
})()
345+
refreshInFlight = refresh
325346

326-
return refreshInFlight
347+
try {
348+
await refresh
349+
} finally {
350+
if (refreshInFlight === refresh) {
351+
refreshInFlight = null
352+
}
353+
}
327354
}
328355

329356
export function resetTinfoilClient(): void {
357+
sessionCacheGeneration++
330358
clientInstance = null
331359
secureClient = null
332360
lastSessionToken = null
333361
cachedSessionToken = null
334362
cachedSessionTokenExpiresAt = null
335363
cachedSessionTokenWasAuthenticated = false
336-
remainingBeforeRequest = null
337364
cachedRateLimit = null
338365
remainingBeforeRequest = null
339366
refreshInFlight = null
340367
}
341368

342369
export function invalidateSessionCache(): void {
370+
sessionCacheGeneration++
343371
cachedSessionToken = null
344372
cachedSessionTokenExpiresAt = null
345373
cachedSessionTokenWasAuthenticated = false
374+
remainingBeforeRequest = null
346375
if (cachedRateLimit !== null) {
347376
cachedRateLimit = null
348377
dispatchRateLimitUpdate()

tests/hooks/use-subscription-status.test.ts

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,11 @@ describe('hasActiveSubscription', () => {
1818
expect(hasActiveSubscription('active', past, now)).toBe(false)
1919
})
2020

21+
it('allows trialing subscriptions with or without a future cutoff', () => {
22+
expect(hasActiveSubscription('trialing', null, now)).toBe(true)
23+
expect(hasActiveSubscription('trialing', future, now)).toBe(true)
24+
})
25+
2126
it('allows canceled subscriptions before their cutoff', () => {
2227
expect(hasActiveSubscription('canceled', future, now)).toBe(true)
2328
})
Lines changed: 108 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,108 @@
1+
import {
2+
getRateLimitInfo,
3+
invalidateSessionCache,
4+
refreshRateLimit,
5+
resetTinfoilClient,
6+
} from '@/services/inference/tinfoil-client'
7+
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
8+
9+
vi.mock('@/config', () => ({
10+
API_BASE_URL: 'https://api.example.com',
11+
DEV_API_KEY: '',
12+
IS_DEV: false,
13+
}))
14+
15+
vi.mock('@/services/auth', () => ({
16+
authTokenManager: {
17+
isInitialized: () => false,
18+
waitForInit: vi.fn(),
19+
getValidToken: vi.fn(),
20+
},
21+
}))
22+
23+
vi.mock('@/utils/error-handling', () => ({
24+
logError: vi.fn(),
25+
}))
26+
27+
const chatKeyResponse = (key: string, remaining: number) =>
28+
new Response(
29+
JSON.stringify({
30+
key,
31+
is_free_tier: true,
32+
rate_limit: {
33+
max_requests: 7,
34+
remaining,
35+
resets_at: '2026-07-24T00:00:00Z',
36+
},
37+
}),
38+
{ status: 200 },
39+
)
40+
41+
describe('tinfoil-client session cache', () => {
42+
beforeEach(() => {
43+
resetTinfoilClient()
44+
localStorage.clear()
45+
})
46+
47+
afterEach(() => {
48+
vi.unstubAllGlobals()
49+
})
50+
51+
it('does not restore stale rate limits after invalidation', async () => {
52+
let resolveStaleResponse: (response: Response) => void = () => {}
53+
const staleResponse = new Promise<Response>((resolve) => {
54+
resolveStaleResponse = resolve
55+
})
56+
const fetchMock = vi
57+
.fn()
58+
.mockReturnValueOnce(staleResponse)
59+
.mockResolvedValueOnce(chatKeyResponse('current-key', 6))
60+
vi.stubGlobal('fetch', fetchMock)
61+
62+
const refresh = refreshRateLimit()
63+
await vi.waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(1))
64+
65+
invalidateSessionCache()
66+
resolveStaleResponse(chatKeyResponse('stale-key', 1))
67+
await refresh
68+
69+
expect(fetchMock).toHaveBeenCalledTimes(2)
70+
expect(getRateLimitInfo()?.remaining).toBe(6)
71+
})
72+
73+
it('keeps tracking a newer refresh when an older refresh finishes', async () => {
74+
let resolveOldResponse: (response: Response) => void = () => {}
75+
let resolveNewResponse: (response: Response) => void = () => {}
76+
const oldResponse = new Promise<Response>((resolve) => {
77+
resolveOldResponse = resolve
78+
})
79+
const newResponse = new Promise<Response>((resolve) => {
80+
resolveNewResponse = resolve
81+
})
82+
const fetchMock = vi
83+
.fn()
84+
.mockReturnValueOnce(oldResponse)
85+
.mockReturnValueOnce(newResponse)
86+
.mockResolvedValueOnce(chatKeyResponse('retried-key', 5))
87+
.mockResolvedValue(chatKeyResponse('unexpected-key', 4))
88+
vi.stubGlobal('fetch', fetchMock)
89+
90+
const oldRefresh = refreshRateLimit()
91+
await vi.waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(1))
92+
93+
resetTinfoilClient()
94+
const newRefresh = refreshRateLimit()
95+
await vi.waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(2))
96+
97+
resolveOldResponse(chatKeyResponse('stale-key', 1))
98+
await oldRefresh
99+
expect(fetchMock).toHaveBeenCalledTimes(3)
100+
101+
const coalescedRefresh = refreshRateLimit()
102+
expect(fetchMock).toHaveBeenCalledTimes(3)
103+
104+
resolveNewResponse(chatKeyResponse('current-key', 6))
105+
await Promise.all([newRefresh, coalescedRefresh])
106+
expect(fetchMock).toHaveBeenCalledTimes(3)
107+
})
108+
})

0 commit comments

Comments
 (0)