diff --git a/apps/web/src/agent/providers/__tests__/gemini.test.ts b/apps/web/src/agent/providers/__tests__/gemini.test.ts index 59558316..691f0840 100644 --- a/apps/web/src/agent/providers/__tests__/gemini.test.ts +++ b/apps/web/src/agent/providers/__tests__/gemini.test.ts @@ -574,4 +574,42 @@ describe("GeminiAdapter.chat()", () => { expect(callArgs.tools).toBeDefined(); expect(callArgs.tools[0].functionDeclarations).toHaveLength(1); }); + + test("retries transient provider failures before returning content", async () => { + const providerError = Object.assign(new Error("Service Unavailable"), { + status: 503, + }); + mockGenerateContent + .mockRejectedValueOnce(providerError) + .mockRejectedValueOnce(providerError) + .mockResolvedValueOnce(makeGeminiResponse({ text: "Recovered" })); + + const adapter = new GeminiAdapter(TEST_CONFIG); + const result = await adapter.chat({ + messages: [makeChatMessage("user", "Hi")], + systemPrompt: "", + tools: [], + }); + + expect(result.content).toBe("Recovered"); + expect(mockGenerateContent).toHaveBeenCalledTimes(3); + }); + + test("does not retry non-transient provider failures", async () => { + const providerError = Object.assign(new Error("Bad request"), { + status: 400, + }); + mockGenerateContent.mockRejectedValueOnce(providerError); + + const adapter = new GeminiAdapter(TEST_CONFIG); + await expect( + adapter.chat({ + messages: [makeChatMessage("user", "Hi")], + systemPrompt: "", + tools: [], + }), + ).rejects.toThrow("Bad request"); + + expect(mockGenerateContent).toHaveBeenCalledTimes(1); + }); }); diff --git a/apps/web/src/agent/providers/gemini.ts b/apps/web/src/agent/providers/gemini.ts index ec1b32c5..5afc2b1e 100644 --- a/apps/web/src/agent/providers/gemini.ts +++ b/apps/web/src/agent/providers/gemini.ts @@ -27,6 +27,51 @@ const TYPE_MAP: Record = { object: SchemaType.OBJECT, }; +const GEMINI_MAX_ATTEMPTS = 5; +const GEMINI_RETRY_BASE_DELAY_MS = process.env.NODE_ENV === "test" ? 0 : 1000; +const GEMINI_RETRYABLE_STATUSES = new Set([429, 500, 502, 503, 504]); + +function getProviderStatus(error: unknown): number | null { + const status = (error as { status?: unknown })?.status; + if (typeof status === "number") return status; + + const message = error instanceof Error ? error.message : String(error); + const match = message.match(/\[(\d{3})\s/); + return match ? Number(match[1]) : null; +} + +function isRetryableProviderError(error: unknown): boolean { + const status = getProviderStatus(error); + return status !== null && GEMINI_RETRYABLE_STATUSES.has(status); +} + +function wait(ms: number): Promise { + return new Promise((resolve) => setTimeout(resolve, ms)); +} + +async function withGeminiRetry(operation: () => Promise): Promise { + let lastError: unknown; + + for (let attempt = 1; attempt <= GEMINI_MAX_ATTEMPTS; attempt++) { + try { + return await operation(); + } catch (error) { + lastError = error; + + if (attempt === GEMINI_MAX_ATTEMPTS || !isRetryableProviderError(error)) { + throw error; + } + + const backoffMs = GEMINI_RETRY_BASE_DELAY_MS * 2 ** (attempt - 1); + const jitterMs = + GEMINI_RETRY_BASE_DELAY_MS === 0 ? 0 : Math.random() * 250; + await wait(backoffMs + jitterMs); + } + } + + throw lastError; +} + /** * Converts internal ChatMessages into Gemini Content[]. * @@ -291,7 +336,9 @@ export class GeminiAdapter implements ProviderAdapter { request.tools = [{ functionDeclarations }]; } - const result = await this.model.generateContent(request); + const result = await withGeminiRetry(() => + this.model.generateContent(request), + ); return fromGeminiResponse(result); } }