diff --git a/src/wrapper/SquidexClient.ts b/src/wrapper/SquidexClient.ts index de3bc42..fb5fdcf 100644 --- a/src/wrapper/SquidexClient.ts +++ b/src/wrapper/SquidexClient.ts @@ -265,11 +265,13 @@ export class SquidexClients { const fetchCore = this.clientOptions.fetcher || fetch; const fetchApi: FetchAPI = async (input, init) => { init ||= {}; + let accessToken: string | undefined; addOptions(init, clientOptions); if (!getHeader(init, "X-AuthRequest")) { - addHeader(init, "Authorization", `Bearer ${await this.tokenApi.getToken()}`); + accessToken = await this.tokenApi.getToken(); + addHeader(init, "Authorization", `Bearer ${accessToken}`); } let response: Response; @@ -290,7 +292,7 @@ export class SquidexClients { if (response && response.status === 401 && !getHeader(init, "X-Retry")) { addHeader(init, "X-Retry", "1"); - this.clearToken(); + this.clearToken(accessToken); return await fetchApi(input, init); } } catch (error: unknown) { @@ -323,8 +325,12 @@ export class SquidexClients { /** * Clears the current token in case it has been expired. */ - clearToken() { - this.tokenStore.clear(); + clearToken(accessToken?: string) { + const token = this.tokenStore.get(); + + if (accessToken === undefined || token?.accessToken === accessToken) { + this.tokenStore.clear(); + } } /** diff --git a/src/wrapper/TokenAPI.ts b/src/wrapper/TokenAPI.ts index 4acf289..5bc9283 100644 --- a/src/wrapper/TokenAPI.ts +++ b/src/wrapper/TokenAPI.ts @@ -14,15 +14,15 @@ export class TokenAPI extends BaseAPI { } public async getToken() { - const promise = (this.tokenPromise ||= (async () => { - const now = new Date().getTime(); - try { - let token = this.tokenStore.get(); + const now = new Date().getTime(); + let token = this.tokenStore.get(); - if (token != null && token.expiresAt > now) { - return token.accessToken; - } + if (token != null && token.expiresAt > now) { + return token.accessToken; + } + const promise = (this.tokenPromise ||= (async () => { + try { const response = await this.request({ path: "/identity-server/connect/token", headers: { diff --git a/tests/tokenRefresh.test.ts b/tests/tokenRefresh.test.ts new file mode 100644 index 0000000..7cd8737 --- /dev/null +++ b/tests/tokenRefresh.test.ts @@ -0,0 +1,139 @@ +import { SquidexClient } from "../src"; + +describe("Token refresh", () => { + it("shares a token request between concurrent API requests", async () => { + let tokenRequestCount = 0; + let resolveTokenRequestStarted: () => void; + let resolvePendingTokenResponse: (response: Response) => void; + + const tokenRequestStarted = new Promise((resolve) => { + resolveTokenRequestStarted = resolve; + }); + + const pendingTokenResponse = new Promise((resolve) => { + resolvePendingTokenResponse = resolve; + }); + + const client = new SquidexClient({ + appName: "my-app", + clientId: "client-id", + clientSecret: "client-secret", + url: "https://squidex.example", + fetcher: async (input) => { + if (input.toString().endsWith("/identity-server/connect/token")) { + tokenRequestCount++; + + if (tokenRequestCount === 1) { + resolveTokenRequestStarted!(); + } + + return pendingTokenResponse; + } + + return new Response(JSON.stringify([])); + }, + }); + + const concurrentApiRequests = [ + client.languages.getLanguages(), + client.languages.getLanguages(), + client.languages.getLanguages(), + ]; + + await tokenRequestStarted; + + expect(tokenRequestCount).toBe(1); + + resolvePendingTokenResponse!(new Response(JSON.stringify({ access_token: "token-1", expires_in: 60_000 }))); + + await Promise.all(concurrentApiRequests); + }); + + it("does not discard a refreshed token after a delayed 401 response", async () => { + let tokenRequestCount = 0; + const tokenOneResponseResolvers: Array<(response: Response) => void> = []; + let resolveInitialApiRequestsReceived: () => void; + + const initialApiRequestsReceived = new Promise((resolve) => { + resolveInitialApiRequestsReceived = resolve; + }); + + const client = new SquidexClient({ + appName: "my-app", + clientId: "client-id", + clientSecret: "client-secret", + url: "https://squidex.example", + fetcher: async (input, init) => { + if (input.toString().endsWith("/identity-server/connect/token")) { + tokenRequestCount++; + return new Response(JSON.stringify({ access_token: `token-${tokenRequestCount}`, expires_in: 60_000 })); + } + + if (new Headers(init?.headers).get("Authorization") === "Bearer token-1") { + return new Promise((resolve) => { + tokenOneResponseResolvers.push(resolve); + + if (tokenOneResponseResolvers.length === 2) { + resolveInitialApiRequestsReceived!(); + } + }); + } + + return new Response(JSON.stringify([])); + }, + }); + + const firstApiRequest = client.languages.getLanguages(); + const secondApiRequest = client.languages.getLanguages(); + + await initialApiRequestsReceived; + + tokenOneResponseResolvers[0](new Response(JSON.stringify({ message: "Unauthorized" }), { status: 401 })); + await firstApiRequest; + + tokenOneResponseResolvers[1](new Response(JSON.stringify({ message: "Unauthorized" }), { status: 401 })); + await secondApiRequest; + + expect(tokenRequestCount).toBe(2); + }); + + it("refreshes a cached token after a 401", async () => { + let tokenRequestCount = 0; + let rejectTokenOne = false; + let tokenOneWasRejected = false; + + const client = new SquidexClient({ + appName: "my-app", + clientId: "client-id", + clientSecret: "client-secret", + url: "https://squidex.example", + fetcher: async (input, init) => { + const url = input.toString(); + const headers = new Headers(init?.headers); + + if (url.endsWith("/identity-server/connect/token")) { + tokenRequestCount++; + + return new Response(JSON.stringify({ access_token: `token-${tokenRequestCount}`, expires_in: 60_000 })); + } + + const authorization = headers.get("Authorization"); + + if (rejectTokenOne && authorization === "Bearer token-1") { + tokenOneWasRejected = true; + return new Response(JSON.stringify({ message: "Unauthorized" }), { status: 401 }); + } + + return new Response(JSON.stringify([])); + }, + }); + + await client.languages.getLanguages(); + + rejectTokenOne = true; + await client.languages.getLanguages(); + + expect(tokenRequestCount).toBe(2); + expect(tokenOneWasRejected).toBe(true); + }); +});