diff --git a/packages/sdk/src/core/agents/retry.test.ts b/packages/sdk/src/core/agents/retry.test.ts index 4b871a4f..38b19d87 100644 --- a/packages/sdk/src/core/agents/retry.test.ts +++ b/packages/sdk/src/core/agents/retry.test.ts @@ -5,6 +5,7 @@ import { Err, Ok } from '~/lib/utils/result.js' import { withLLMRetry, withRetry } from './retry.js' const timeout: LLMError = { type: 'timeout', message: 'Request timed out' } +const serverError: LLMError = { type: 'server_error', message: 'Bad gateway' } describe('withRetry abort handling', () => { test('a cancel after a retryable failure reports the cancel, not the failure', async () => { @@ -54,14 +55,14 @@ describe('withRetry abort handling', () => { const result = await withLLMRetry( async () => { attempts++ - return Err(timeout) + return Err(serverError) }, { maxAttempts: 3, baseDelayMs: 1, maxDelayMs: 1 }, ) expect(attempts).toBe(3) expect(result.ok).toBe(false) - if (!result.ok) expect(result.error.type).toBe('timeout') + if (!result.ok) expect(result.error.type).toBe('server_error') }) test('the final failed attempt neither calculates nor logs another retry', async () => { @@ -148,3 +149,35 @@ describe('withRetry abort handling', () => { expect(result).toEqual(Ok('done')) }) }) + +describe('withLLMRetry attempt caps', () => { + test('a timeout is attempted twice', async () => { + let attempts = 0 + + const result = await withLLMRetry( + async () => { + attempts++ + return Err(timeout) + }, + { baseDelayMs: 1, maxDelayMs: 1 }, + ) + + expect(attempts).toBe(2) + expect(result).toEqual(Err(timeout)) + }) + + test('other retryable errors keep the full attempt budget', async () => { + let attempts = 0 + + const result = await withLLMRetry( + async () => { + attempts++ + return Err(serverError) + }, + { baseDelayMs: 1, maxDelayMs: 1 }, + ) + + expect(attempts).toBe(5) + expect(result).toEqual(Err(serverError)) + }) +}) diff --git a/packages/sdk/src/core/agents/retry.ts b/packages/sdk/src/core/agents/retry.ts index 0f1a9004..d36747d2 100644 --- a/packages/sdk/src/core/agents/retry.ts +++ b/packages/sdk/src/core/agents/retry.ts @@ -27,6 +27,8 @@ export const DEFAULT_RETRY_OPTIONS: Required = { export interface WithRetryOptions extends RetryOptions { isRetryable: (error: E) => boolean getRetryDelay?: (error: E) => number | undefined + /** Lower attempt cap for the error just returned; `undefined` keeps `maxAttempts`. */ + getMaxAttempts?: (error: E) => number | undefined /** Error to return when the caller's signal aborts, whichever attempt that happens on. */ abortError?: E logger?: Logger @@ -83,7 +85,8 @@ export async function withRetry( return Err(options.abortError ?? lastError) } - if (attempt >= opts.maxAttempts || !options.isRetryable(lastError)) { + const maxAttempts = Math.min(opts.maxAttempts, options.getMaxAttempts?.(lastError) ?? opts.maxAttempts) + if (attempt >= maxAttempts || !options.isRetryable(lastError)) { return result } @@ -139,6 +142,16 @@ export function isRetryableLLMError(error: LLMError): boolean { ) } +/** + * A timed-out request already cost a full timeout, and a response too long to finish + * in time times out again on the identical retry. + */ +const MAX_LLM_TIMEOUT_ATTEMPTS = 2 + +function getLLMMaxAttempts(error: LLMError): number | undefined { + return error.type === 'timeout' ? MAX_LLM_TIMEOUT_ATTEMPTS : undefined +} + /** * Gets retry delay from LLM error if available (e.g., rate limit retry-after). */ @@ -157,6 +170,7 @@ export async function withLLMRetry( ...options, isRetryable: isRetryableLLMError, getRetryDelay: getLLMRetryDelay, + getMaxAttempts: getLLMMaxAttempts, abortError: { type: 'aborted', message: 'Request was aborted' }, context: 'LLM inference', }) diff --git a/packages/sdk/src/core/llm/anthropic.ts b/packages/sdk/src/core/llm/anthropic.ts index 31b18776..e43f897a 100644 --- a/packages/sdk/src/core/llm/anthropic.ts +++ b/packages/sdk/src/core/llm/anthropic.ts @@ -18,7 +18,7 @@ import type { RawToolSpec, } from './provider.js' import { mapProviderError } from './provider.js' -import { ProviderRequestAbortError, runProviderRequest } from './provider-request.js' +import { DEFAULT_PROVIDER_REQUEST_TIMEOUT_MS, ProviderRequestAbortError, runProviderRequest } from './provider-request.js' import { sanitizeProviderMessages } from './message-sanitization.js' import type { RoutableLLMProvider } from './routing-provider.js' @@ -252,7 +252,7 @@ export class AnthropicProvider implements RoutableLLMProvider { this.logger = config.logger this.imageProcessor = config.imageProcessor this.thinkingBudget = config.thinkingBudget - this.timeout = config.timeout ?? 120000 + this.timeout = config.timeout ?? DEFAULT_PROVIDER_REQUEST_TIMEOUT_MS this.fetchFn = config.fetch ?? globalThis.fetch } diff --git a/packages/sdk/src/core/llm/openrouter.ts b/packages/sdk/src/core/llm/openrouter.ts index 62ddf0d4..8043ed05 100644 --- a/packages/sdk/src/core/llm/openrouter.ts +++ b/packages/sdk/src/core/llm/openrouter.ts @@ -19,7 +19,7 @@ import type { RawInferenceRequest, } from './provider.js' import { mapProviderError } from './provider.js' -import { ProviderRequestAbortError, runProviderRequest } from './provider-request.js' +import { DEFAULT_PROVIDER_REQUEST_TIMEOUT_MS, ProviderRequestAbortError, runProviderRequest } from './provider-request.js' import { sanitizeProviderMessages } from './message-sanitization.js' // ============================================================================ @@ -197,7 +197,7 @@ export class OpenRouterProvider implements LLMProvider { this.defaultModel = config.defaultModel ?? 'anthropic/claude-sonnet-4.5' this.logger = config.logger this.imageProcessor = config.imageProcessor - this.timeout = config.timeout ?? 120000 + this.timeout = config.timeout ?? DEFAULT_PROVIDER_REQUEST_TIMEOUT_MS this.fetchFn = config.fetch ?? globalThis.fetch } diff --git a/packages/sdk/src/core/llm/provider-request.ts b/packages/sdk/src/core/llm/provider-request.ts index 8ed9df34..3f4737b6 100644 --- a/packages/sdk/src/core/llm/provider-request.ts +++ b/packages/sdk/src/core/llm/provider-request.ts @@ -1,5 +1,8 @@ export type ProviderRequestAbortCause = 'caller' | 'timeout' +/** Bounds the whole non-streamed response, so it caps how many output tokens one call can produce. */ +export const DEFAULT_PROVIDER_REQUEST_TIMEOUT_MS = 300_000 + export class ProviderRequestAbortError extends Error { constructor(readonly abortCause: ProviderRequestAbortCause) { super(abortCause === 'caller' ? 'Request was aborted' : 'Request timed out')