diff --git a/src/core.ts b/src/core.ts index 10b9d3f..946177b 100644 --- a/src/core.ts +++ b/src/core.ts @@ -625,7 +625,7 @@ export abstract class APIClient { const maxRetries = options.maxRetries ?? this.maxRetries; timeoutMillis = this.calculateDefaultRetryTimeoutMillis(retriesRemaining, maxRetries); } - await sleep(timeoutMillis); + await sleepWithAbort(timeoutMillis, options.signal); return this.makeRequest(options, retriesRemaining - 1); } @@ -1019,6 +1019,28 @@ const isAbsoluteURL = (url: string): boolean => { export const sleep = (ms: number) => new Promise((resolve) => setTimeout(resolve, ms)); +const sleepWithAbort = (ms: number, signal?: AbortSignal | null): Promise => { + return new Promise((resolve, reject) => { + if (signal?.aborted) { + reject(new APIUserAbortError()); + return; + } + + const abortHandler = () => { + clearTimeout(timeout); + signal?.removeEventListener('abort', abortHandler); + reject(new APIUserAbortError()); + }; + const timeout = setTimeout(() => { + signal?.removeEventListener('abort', abortHandler); + resolve(); + }, ms); + + signal?.addEventListener('abort', abortHandler); + if (signal?.aborted) abortHandler(); + }); +}; + const validatePositiveInteger = (name: string, n: unknown): number => { if (typeof n !== 'number' || !Number.isInteger(n)) { throw new BrowserbaseError(`${name} must be an integer`); diff --git a/tests/index.test.ts b/tests/index.test.ts index 7639b23..9685914 100644 --- a/tests/index.test.ts +++ b/tests/index.test.ts @@ -265,6 +265,54 @@ describe('request building', () => { }); describe('retries', () => { + test('custom signal aborts retry backoff promptly', async () => { + let count = 0; + const testFetch = async (): Promise => { + count++; + return new Response(undefined, { + status: 429, + headers: { 'Retry-After-Ms': '1000' }, + }); + }; + const client = new Browserbase({ apiKey: 'My API Key', fetch: testFetch, maxRetries: 1 }); + const controller = new AbortController(); + const addEventListener = jest.spyOn(controller.signal, 'addEventListener'); + const removeEventListener = jest.spyOn(controller.signal, 'removeEventListener'); + const request = client.get('/foo', { signal: controller.signal }); + + while (addEventListener.mock.calls.length < 2) await Promise.resolve(); + const retryAbortHandler = addEventListener.mock.calls[1]![1]; + controller.abort(); + + await expect( + Promise.race([ + request, + new Promise((_, reject) => setTimeout(() => reject(new Error('abort was not prompt')), 100)), + ]), + ).rejects.toThrow(APIUserAbortError); + expect(count).toBe(1); + expect(removeEventListener).toHaveBeenCalledWith('abort', retryAbortHandler); + }); + + test('cleans up the abort listener after retry backoff completes', async () => { + let count = 0; + const testFetch = async (): Promise => { + if (count++ === 0) { + return new Response(undefined, { status: 429, headers: { 'Retry-After-Ms': '10' } }); + } + return new Response(JSON.stringify({ a: 1 }), { headers: { 'Content-Type': 'application/json' } }); + }; + const client = new Browserbase({ apiKey: 'My API Key', fetch: testFetch, maxRetries: 1 }); + const controller = new AbortController(); + const addEventListener = jest.spyOn(controller.signal, 'addEventListener'); + const removeEventListener = jest.spyOn(controller.signal, 'removeEventListener'); + + await expect(client.get('/foo', { signal: controller.signal })).resolves.toEqual({ a: 1 }); + + const retryAbortHandler = addEventListener.mock.calls[1]![1]; + expect(removeEventListener).toHaveBeenCalledWith('abort', retryAbortHandler); + }); + test('retry on timeout', async () => { let count = 0; const testFetch = async (url: RequestInfo, { signal }: RequestInit = {}): Promise => {