Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
73 changes: 73 additions & 0 deletions src/main/commitMessageProviders.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,79 @@ afterEach(() => {
});

describe("commit message provider cancellation", () => {
it.each(apiProviderFactories)("keeps $name active while response bytes arrive for more than a minute", async ({ create }) => {
vi.useFakeTimers();
let body!: ReadableStreamDefaultController<Uint8Array>;
const response = new Response(new ReadableStream<Uint8Array>({
start(controller) { body = controller; }
}));
const provider = create(async () => response);
const generation = provider.generate(providerInput);
const result = generation.catch((error: unknown) => error);
const encoder = new TextEncoder();
for (let index = 0; index < 3; index += 1) {
await vi.advanceTimersByTimeAsync(30_000);
body.enqueue(encoder.encode(" "));
}
body.enqueue(encoder.encode(JSON.stringify({
choices: [{ message: { content: "feat: café" } }],
output_text: "feat: café",
content: [{ type: "text", text: "feat: café" }]
})));
body.close();
expect(await result).toMatchObject({ text: "feat: café" });
expect(vi.getTimerCount()).toBe(0);
});

it.each(apiProviderFactories)("times out $name after response activity stops", async ({ create }) => {
vi.useFakeTimers();
let body!: ReadableStreamDefaultController<Uint8Array>;
const cancel = vi.fn();
const response = new Response(new ReadableStream<Uint8Array>({
start(controller) { body = controller; },
cancel
}));
const provider = create(async () => response);
const settled = vi.fn();
const result = provider.generate(providerInput).catch((error: unknown) => error).then((value) => {
settled();
return value;
});
await vi.advanceTimersByTimeAsync(30_000);
body.enqueue(new TextEncoder().encode(" "));
await vi.advanceTimersByTimeAsync(59_999);
expect(settled).not.toHaveBeenCalled();
// Empty chunks do not indicate response progress.
body.enqueue(new Uint8Array());
await vi.advanceTimersByTimeAsync(1);
expect(await result).toMatchObject({ name: "TimeoutError" });
expect(cancel).toHaveBeenCalledOnce();
expect(response.body?.locked).toBe(false);
expect(vi.getTimerCount()).toBe(0);
});

it.each(apiProviderFactories)("cancels $name after response activity extends the request", async ({ create }) => {
vi.useFakeTimers();
const controller = new AbortController();
let body!: ReadableStreamDefaultController<Uint8Array>;
const cancel = vi.fn();
const response = new Response(new ReadableStream<Uint8Array>({
start(streamController) { body = streamController; },
cancel
}));
const provider = create(async () => response);
const result = provider.generate({ ...providerInput, signal: controller.signal }).catch((error: unknown) => error);
await vi.advanceTimersByTimeAsync(30_000);
body.enqueue(new TextEncoder().encode(" "));
await vi.advanceTimersByTimeAsync(40_000);
const reason = new DOMException("Generation cancelled.", "AbortError");
controller.abort(reason);
expect(await result).toBe(reason);
expect(cancel).toHaveBeenCalledWith(reason);
expect(response.body?.locked).toBe(false);
expect(vi.getTimerCount()).toBe(0);
});

it("propagates an external abort to an in-flight API request", async () => {
const controller = new AbortController();
let requestStarted!: () => void;
Expand Down
3 changes: 3 additions & 0 deletions src/main/commitMessageProviders.ts
Original file line number Diff line number Diff line change
Expand Up @@ -163,6 +163,7 @@ export class OpenRouterCommitMessageProvider implements CommitMessageProvider {
})
}, {
timeoutMs: API_TIMEOUT_MS,
timeoutMode: "inactivity",
timeoutReason: createApiTimeoutReason(),
...(input.signal ? { signal: input.signal } : {}),
createJsonFallback: () => ({} as OpenRouterResponse)
Expand Down Expand Up @@ -193,6 +194,7 @@ export class OpenAiCommitMessageProvider implements CommitMessageProvider {
})
}, {
timeoutMs: API_TIMEOUT_MS,
timeoutMode: "inactivity",
timeoutReason: createApiTimeoutReason(),
...(input.signal ? { signal: input.signal } : {}),
createJsonFallback: () => ({} as OpenAiResponse)
Expand Down Expand Up @@ -242,6 +244,7 @@ export class AnthropicCommitMessageProvider implements CommitMessageProvider {
})
}, {
timeoutMs: API_TIMEOUT_MS,
timeoutMode: "inactivity",
timeoutReason: createApiTimeoutReason(),
...(input.signal ? { signal: input.signal } : {}),
createJsonFallback: () => ({} as AnthropicResponse)
Expand Down
34 changes: 31 additions & 3 deletions src/main/fetchJsonWithTimeout.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@ type Fetch = typeof fetch;

export interface FetchJsonWithTimeoutOptions<T> {
timeoutMs: number;
/** Inactivity mode resets the timeout on headers and non-empty body chunks. */
timeoutMode?: "total" | "inactivity";
signal?: AbortSignal;
timeoutReason?: unknown;
createJsonFallback?: () => T;
Expand Down Expand Up @@ -47,7 +49,9 @@ export async function fetchJsonWithTimeout<T>(
...init,
signal: controller.signal
}), controller.signal);
const payload = await parseJson(response, controller.signal, options.createJsonFallback);
const onActivity = options.timeoutMode === "inactivity" ? () => { timeout.refresh(); } : undefined;
onActivity?.();
const payload = await parseJson(response, controller.signal, options.createJsonFallback, onActivity);
return { response, payload };
} catch (error) {
if (controller.signal.aborted) {
Expand All @@ -65,10 +69,34 @@ export async function fetchJsonWithTimeout<T>(
async function parseJson<T>(
response: Response,
signal: AbortSignal,
createFallback: (() => T) | undefined
createFallback: (() => T) | undefined,
onActivity: (() => void) | undefined
): Promise<T> {
try {
return await raceWithAbort(response.json() as Promise<T>, signal);
if (!onActivity || !response.body) {
return await raceWithAbort(response.json() as Promise<T>, signal);
}
const reader = response.body.getReader();
const decoder = new TextDecoder();
const chunks: string[] = [];
try {
while (true) {
const { done, value } = await raceWithAbort(reader.read(), signal);
if (done) break;
if (value.byteLength > 0) {
onActivity();
chunks.push(decoder.decode(value, { stream: true }));
}
}
chunks.push(decoder.decode());
return JSON.parse(chunks.join("")) as T;
} finally {
if (signal.aborted) {
// Cancellation must not wait for an unresponsive underlying stream.
void reader.cancel(signal.reason).catch(() => undefined);
}
reader.releaseLock();
}
} catch (error) {
if (signal.aborted) {
throw signal.reason;
Expand Down