diff --git a/README.md b/README.md index b9711948e..83b23504f 100644 --- a/README.md +++ b/README.md @@ -885,16 +885,33 @@ theme: dark: titanium light: light +enabledModels: + - "anthropic/*" + - "*gpt*" + - "gemini-2.5-pro:high" + modelRoles: default: anthropic/claude-sonnet-4-20250514 plan: anthropic/claude-opus-4-1:high smol: anthropic/claude-sonnet-4-20250514 defaultThinkingLevel: high -enabledModels: - - anthropic/* - - "*gpt*" - - gemini-2.5-pro:high +retry: + enabled: true + # Number of retries before giving up on rate limits/server errors + maxRetries: 3 + # Wait this long as a base (exponentially backed off) unless the API provides a retry-after-ms + baseDelayMs: 2000 + # Configure role-specific model fallback chains + fallbackChains: + default: + - "openai/gpt-4o-mini" + - "openai/gpt-4o" + plan: + - "anthropic/claude-sonnet-4-6:high" + - "openai/o3:high" + # Whether to revert to the primary model when a fallback's cooldown expires + fallbackRevertPolicy: cooldown-expiry steeringMode: one-at-a-time followUpMode: one-at-a-time interruptMode: immediate @@ -914,10 +931,6 @@ compaction: skills: enabled: true -retry: - enabled: true - maxRetries: 3 - baseDelayMs: 2000 terminal: showImages: true diff --git a/packages/coding-agent/src/config/model-registry.ts b/packages/coding-agent/src/config/model-registry.ts index e17679a7f..ff85de248 100644 --- a/packages/coding-agent/src/config/model-registry.ts +++ b/packages/coding-agent/src/config/model-registry.ts @@ -28,6 +28,7 @@ import { import { isRecord, logger } from "@oh-my-pi/pi-utils"; import { type Static, Type } from "@sinclair/typebox"; import { type ConfigError, ConfigFile } from "../config"; +import { parseModelString } from "../config/model-resolver"; import { isValidThemeColor, type ThemeColor } from "../modes/theme/theme"; import type { AuthStorage, OAuthCredential } from "../session/auth-storage"; import type { Settings } from "./settings"; @@ -721,6 +722,14 @@ function buildCustomModel( return finalizeCustomModel(model, options); } +function normalizeSuppressedSelector(selector: string): string { + const trimmed = selector.trim(); + if (!trimmed) return trimmed; + const parsed = parseModelString(trimmed); + if (!parsed) return trimmed; + return `${parsed.provider}/${parsed.id}`; +} + /** * Model registry - loads and manages models, resolves API keys via AuthStorage. */ @@ -737,6 +746,7 @@ export class ModelRegistry { #registeredProviderSources: Set = new Set(); #providerDiscoveryStates: Map = new Map(); #cacheDbPath?: string; + #suppressedSelectors: Map = new Map(); #backgroundRefresh?: Promise; #lastDiscoveryWarnings: Map = new Map(); @@ -766,6 +776,7 @@ export class ModelRegistry { */ async refresh(strategy: ModelRefreshStrategy = "online-if-uncached"): Promise { this.#reloadStaticModels(); + this.#suppressedSelectors.clear(); await this.#refreshRuntimeDiscoveries(strategy); } @@ -789,6 +800,11 @@ export class ModelRegistry { async refreshProvider(providerId: string, strategy: ModelRefreshStrategy = "online"): Promise { this.#reloadStaticModels(); + for (const selector of this.#suppressedSelectors.keys()) { + if (selector.startsWith(`${providerId}/`)) { + this.#suppressedSelectors.delete(selector); + } + } await this.#refreshRuntimeDiscoveries(strategy, new Set([providerId])); } @@ -1825,6 +1841,27 @@ export class ModelRegistry { }); } } + + /** + * Suppress a specific model selector (e.g., "provider/id") until a specific timestamp. + */ + suppressSelector(selector: string, untilMs: number): void { + this.#suppressedSelectors.set(normalizeSuppressedSelector(selector), untilMs); + } + + /** + * Check if a model selector is currently suppressed due to rate limits. + */ + isSelectorSuppressed(selector: string): boolean { + const normalizedSelector = normalizeSuppressedSelector(selector); + const suppressedUntil = this.#suppressedSelectors.get(normalizedSelector); + if (!suppressedUntil) return false; + if (suppressedUntil <= Date.now()) { + this.#suppressedSelectors.delete(normalizedSelector); + return false; + } + return true; + } } /** diff --git a/packages/coding-agent/src/config/settings-schema.ts b/packages/coding-agent/src/config/settings-schema.ts index 17ecfcf44..865b5783e 100644 --- a/packages/coding-agent/src/config/settings-schema.ts +++ b/packages/coding-agent/src/config/settings-schema.ts @@ -525,6 +525,18 @@ export const SETTINGS_SCHEMA = { }, "retry.baseDelayMs": { type: "number", default: 2000 }, + "retry.fallbackChains": { type: "record", default: {} as Record }, + "retry.fallbackRevertPolicy": { + type: "enum", + values: ["cooldown-expiry", "never"] as const, + default: "cooldown-expiry", + ui: { + tab: "model", + label: "Fallback Revert Policy", + description: "When to return to the primary model after a fallback", + submenu: true, + }, + }, // ──────────────────────────────────────────────────────────────────────── // Interaction diff --git a/packages/coding-agent/src/modes/components/settings-defs.ts b/packages/coding-agent/src/modes/components/settings-defs.ts index 4c474448a..bd5a1de37 100644 --- a/packages/coding-agent/src/modes/components/settings-defs.ts +++ b/packages/coding-agent/src/modes/components/settings-defs.ts @@ -117,6 +117,15 @@ const OPTION_PROVIDERS: Partial> = { { value: "5", label: "5 retries" }, { value: "10", label: "10 retries" }, ], + // Retry fallback revert policy + "retry.fallbackRevertPolicy": [ + { + value: "cooldown-expiry", + label: "Cooldown expiry", + description: "Return to the primary model after its suppression window ends", + }, + { value: "never", label: "Never", description: "Stay on the fallback model until manually changed" }, + ], // Task max concurrency "task.maxConcurrency": [ { value: "0", label: "Unlimited" }, diff --git a/packages/coding-agent/src/modes/controllers/event-controller.ts b/packages/coding-agent/src/modes/controllers/event-controller.ts index d7085930c..7e18e4986 100644 --- a/packages/coding-agent/src/modes/controllers/event-controller.ts +++ b/packages/coding-agent/src/modes/controllers/event-controller.ts @@ -534,6 +534,16 @@ export class EventController { break; } + case "retry_fallback_applied": { + this.ctx.showWarning(`Fallback: ${event.from} -> ${event.to}`); + break; + } + + case "retry_fallback_succeeded": { + this.ctx.showStatus(`Fallback succeeded on ${event.model}`); + break; + } + case "ttsr_triggered": { const component = new TtsrNotificationComponent(event.rules); component.setExpanded(this.ctx.toolOutputExpanded); diff --git a/packages/coding-agent/src/modes/interactive-mode.ts b/packages/coding-agent/src/modes/interactive-mode.ts index caf1d8621..66acf5ebc 100644 --- a/packages/coding-agent/src/modes/interactive-mode.ts +++ b/packages/coding-agent/src/modes/interactive-mode.ts @@ -306,6 +306,11 @@ export class InteractiveMode implements InteractiveModeContext { const startupQuiet = settings.get("startup.quiet"); + for (const warning of this.session.configWarnings) { + this.ui.addChild(new Text(theme.fg("warning", `Warning: ${warning}`), 1, 0)); + this.ui.addChild(new Spacer(1)); + } + if (!startupQuiet) { // Add welcome header const welcome = new WelcomeComponent(this.#version, modelName, providerName, recentSessions, lspServerInfo); diff --git a/packages/coding-agent/src/session/agent-session.ts b/packages/coding-agent/src/session/agent-session.ts index d9e4737d9..c6bf87d62 100644 --- a/packages/coding-agent/src/session/agent-session.ts +++ b/packages/coding-agent/src/session/agent-session.ts @@ -55,7 +55,12 @@ import { abortableSleep, getAgentDbPath, isEnoent, logger } from "@oh-my-pi/pi-u import type { AsyncJob, AsyncJobManager } from "../async"; import type { Rule } from "../capability/rule"; import { MODEL_ROLE_IDS, type ModelRegistry } from "../config/model-registry"; -import { extractExplicitThinkingSelector, parseModelString, resolveModelRoleValue } from "../config/model-resolver"; +import { + extractExplicitThinkingSelector, + formatModelString, + parseModelString, + resolveModelRoleValue, +} from "../config/model-resolver"; import { expandPromptTemplate, type PromptTemplate, renderPromptTemplate } from "../config/prompt-templates"; import type { Settings, SkillsSettings } from "../config/settings"; import { type BashResult, executeBash as executeBashCommand } from "../exec/bash-executor"; @@ -170,6 +175,8 @@ export type AgentSessionEvent = } | { type: "auto_retry_start"; attempt: number; maxAttempts: number; delayMs: number; errorMessage: string } | { type: "auto_retry_end"; success: boolean; attempt: number; finalError?: string } + | { type: "retry_fallback_applied"; from: string; to: string; role: string } + | { type: "retry_fallback_succeeded"; model: string; role: string } | { type: "ttsr_triggered"; rules: Rule[] } | { type: "todo_reminder"; todos: TodoItem[]; attempt: number; maxAttempts: number } | { type: "todo_auto_clear" }; @@ -315,6 +322,46 @@ interface HandoffOptions { const AUTO_HANDOFF_THRESHOLD_FOCUS = renderPromptTemplate(autoHandoffThresholdFocusPrompt); +type RetryFallbackChains = Record; + +type RetryFallbackRevertPolicy = "never" | "cooldown-expiry"; + +interface RetryFallbackSelector { + raw: string; + provider: string; + id: string; + thinkingLevel: ThinkingLevel | undefined; +} + +interface ActiveRetryFallbackState { + role: string; + originalSelector: string; + originalThinkingLevel: ThinkingLevel | undefined; + lastAppliedFallbackThinkingLevel: ThinkingLevel | undefined; +} + +function parseRetryFallbackSelector(selector: string): RetryFallbackSelector | undefined { + const trimmed = selector.trim(); + if (!trimmed) return undefined; + const parsed = parseModelString(trimmed); + if (!parsed) return undefined; + return { + raw: trimmed, + provider: parsed.provider, + id: parsed.id, + thinkingLevel: parsed.thinkingLevel, + }; +} + +function formatRetryFallbackSelector(model: Model, thinkingLevel: ThinkingLevel | undefined): string { + const selector = formatModelString(model); + return thinkingLevel ? `${selector}:${thinkingLevel}` : selector; +} + +function formatRetryFallbackBaseSelector(selector: RetryFallbackSelector): string { + return `${selector.provider}/${selector.id}`; +} + const noOpUIContext: ExtensionUIContext = { select: async (_title, _options, _dialogOptions) => undefined, confirm: async (_title, _message, _dialogOptions) => false, @@ -352,6 +399,7 @@ export class AgentSession { readonly sessionManager: SessionManager; readonly settings: Settings; readonly searchDb: SearchDb | undefined; + readonly configWarnings: string[] = []; #asyncJobManager: AsyncJobManager | undefined = undefined; #scopedModels: Array<{ model: Model; thinkingLevel?: ThinkingLevel }>; @@ -391,7 +439,7 @@ export class AgentSession { #retryAttempt = 0; #retryPromise: Promise | undefined = undefined; #retryResolve: (() => void) | undefined = undefined; - + #activeRetryFallback: ActiveRetryFallbackState | undefined = undefined; // Todo completion reminder state #todoReminderCount = 0; #todoPhases: TodoPhase[] = []; @@ -478,6 +526,7 @@ export class AgentSession { this.#customCommands = config.customCommands ?? []; this.#skillsSettings = config.skillsSettings; this.#modelRegistry = config.modelRegistry; + this.#validateRetryFallbackChains(); this.#toolRegistry = config.toolRegistry ?? new Map(); this.#transformContext = config.transformContext ?? (messages => messages); this.#onPayload = config.onPayload; @@ -787,6 +836,13 @@ export class AgentSession { assistantMsg.stopReason !== "aborted" && this.#retryAttempt > 0 ) { + if (this.#activeRetryFallback && this.model) { + await this.#emitSessionEvent({ + type: "retry_fallback_succeeded", + model: formatRetryFallbackSelector(this.model, this.thinkingLevel), + role: this.#activeRetryFallback.role, + }); + } await this.#emitSessionEvent({ type: "auto_retry_end", success: true, @@ -985,6 +1041,7 @@ export class AgentSession { return; } try { + await this.#maybeRestoreRetryFallbackPrimary(); await this.agent.continue(); } catch { options?.onError?.(); @@ -2288,6 +2345,8 @@ export class AgentSession { // Reset todo reminder count on new user prompt this.#todoReminderCount = 0; + await this.#maybeRestoreRetryFallbackPrimary(); + // Validate model if (!this.model) { throw new Error( @@ -3108,6 +3167,7 @@ export class AgentSession { throw new Error(`No API key for ${model.provider}/${model.id}`); } + this.#clearActiveRetryFallback(); this.#setModelWithProviderSessionReset(model); this.sessionManager.appendModelChange(`${model.provider}/${model.id}`, role); this.settings.setModelRole(role, this.#formatRoleModelValue(role, model)); @@ -3128,6 +3188,7 @@ export class AgentSession { throw new Error(`No API key for ${model.provider}/${model.id}`); } + this.#clearActiveRetryFallback(); this.#setModelWithProviderSessionReset(model); this.sessionManager.appendModelChange(`${model.provider}/${model.id}`, "temporary"); this.settings.getStorage()?.recordModelUsage(`${model.provider}/${model.id}`); @@ -3253,6 +3314,7 @@ export class AgentSession { const next = scopedModels[nextIndex]; // Apply model + this.#clearActiveRetryFallback(); this.#setModelWithProviderSessionReset(next.model); this.sessionManager.appendModelChange(`${next.model.provider}/${next.model.id}`); this.settings.setModelRole("default", this.#formatRoleModelValue("default", next.model)); @@ -3281,11 +3343,11 @@ export class AgentSession { throw new Error(`No API key for ${nextModel.provider}/${nextModel.id}`); } + this.#clearActiveRetryFallback(); this.#setModelWithProviderSessionReset(nextModel); this.sessionManager.appendModelChange(`${nextModel.provider}/${nextModel.id}`); this.settings.setModelRole("default", this.#formatRoleModelValue("default", nextModel)); this.settings.getStorage()?.recordModelUsage(`${nextModel.provider}/${nextModel.id}`); - // Re-apply the current thinking level for the newly selected model this.setThinkingLevel(this.thinkingLevel); @@ -4768,6 +4830,217 @@ export class AgentSession { ); } + #getRetryFallbackChains(): RetryFallbackChains { + const configuredChains = this.settings.get("retry.fallbackChains"); + if (!configuredChains || typeof configuredChains !== "object") return {}; + return configuredChains as RetryFallbackChains; + } + + #validateRetryFallbackChains(): void { + const configuredChains = this.settings.get("retry.fallbackChains"); + if (configuredChains === undefined) return; + if (!configuredChains || typeof configuredChains !== "object" || Array.isArray(configuredChains)) { + const msg = "retry.fallbackChains must be a mapping of role names to selector arrays."; + logger.warn(msg); + this.configWarnings.push(msg); + return; + } + + for (const [role, chain] of Object.entries(configuredChains)) { + if (!Array.isArray(chain)) { + const msg = `Fallback chain for role '${role}' must be an array of selector strings.`; + logger.warn(msg); + this.configWarnings.push(msg); + continue; + } + for (const selectorStr of chain) { + if (typeof selectorStr !== "string") { + const msg = `Fallback chain for role '${role}' contains a non-string selector.`; + logger.warn(msg); + this.configWarnings.push(msg); + continue; + } + const parsed = parseRetryFallbackSelector(selectorStr); + if (!parsed) { + const msg = `Invalid fallback selector format in role '${role}': ${selectorStr}`; + logger.warn(msg); + this.configWarnings.push(msg); + continue; + } + const exists = this.#modelRegistry.find(parsed.provider, parsed.id); + if (!exists) { + const msg = `Fallback chain for role '${role}' references unknown model: ${selectorStr}`; + logger.warn(msg); + this.configWarnings.push(msg); + } + } + } + } + + #getRetryFallbackRevertPolicy(): RetryFallbackRevertPolicy { + return this.settings.get("retry.fallbackRevertPolicy") === "never" ? "never" : "cooldown-expiry"; + } + + #getRetryFallbackPrimarySelector(role: string): RetryFallbackSelector | undefined { + const configuredSelector = this.settings.getModelRole(role); + return configuredSelector ? parseRetryFallbackSelector(configuredSelector) : undefined; + } + + #clearActiveRetryFallback(): void { + this.#activeRetryFallback = undefined; + } + + #isRetryFallbackSelectorSuppressed(selector: RetryFallbackSelector): boolean { + return this.#modelRegistry.isSelectorSuppressed(selector.raw); + } + + #noteRetryFallbackCooldown(currentSelector: string, retryAfterMs: number | undefined, errorMessage: string): void { + let cooldownMs = retryAfterMs; + if (!cooldownMs || cooldownMs <= 0) { + const reason = parseRateLimitReason(errorMessage); + cooldownMs = reason === "UNKNOWN" ? 5 * 60 * 1000 : calculateRateLimitBackoffMs(reason); + } + this.#modelRegistry.suppressSelector(currentSelector, Date.now() + cooldownMs); + } + + #resolveRetryFallbackRole(currentSelector: string): string | undefined { + const parsedCurrent = parseRetryFallbackSelector(currentSelector); + if (!parsedCurrent) return undefined; + const currentBaseSelector = formatRetryFallbackBaseSelector(parsedCurrent); + for (const role of Object.keys(this.#getRetryFallbackChains())) { + const primarySelector = this.#getRetryFallbackPrimarySelector(role); + if (!primarySelector) continue; + if (primarySelector.raw === currentSelector) return role; + if (formatRetryFallbackBaseSelector(primarySelector) === currentBaseSelector) return role; + } + return undefined; + } + + #getRetryFallbackEffectiveChain(role: string): RetryFallbackSelector[] { + const primarySelector = this.#getRetryFallbackPrimarySelector(role); + if (!primarySelector) return []; + const chain = [primarySelector]; + const seen = new Set([primarySelector.raw]); + for (const selector of this.#getRetryFallbackChains()[role] ?? []) { + const parsed = parseRetryFallbackSelector(selector); + if (!parsed || seen.has(parsed.raw)) continue; + seen.add(parsed.raw); + chain.push(parsed); + } + return chain; + } + + #findRetryFallbackCandidates(role: string, currentSelector: string): RetryFallbackSelector[] { + const chain = this.#getRetryFallbackEffectiveChain(role); + if (chain.length <= 1) return []; + const parsedCurrent = parseRetryFallbackSelector(currentSelector); + const currentBaseSelector = parsedCurrent ? formatRetryFallbackBaseSelector(parsedCurrent) : undefined; + const exactIndex = chain.findIndex(selector => selector.raw === currentSelector); + if (exactIndex >= 0) return chain.slice(exactIndex + 1); + const baseIndex = currentBaseSelector + ? chain.findIndex(selector => formatRetryFallbackBaseSelector(selector) === currentBaseSelector) + : -1; + if (baseIndex >= 0) return chain.slice(baseIndex + 1); + return chain.slice(1); + } + + async #applyRetryFallbackCandidate( + role: string, + selector: RetryFallbackSelector, + currentSelector: string, + ): Promise { + const candidate = this.#modelRegistry.find(selector.provider, selector.id); + if (!candidate) { + throw new Error(`Retry fallback model not found: ${selector.raw}`); + } + const apiKey = await this.#modelRegistry.getApiKey(candidate, this.sessionId); + if (!apiKey) { + throw new Error(`No API key for retry fallback ${selector.raw}`); + } + + const currentThinkingLevel = this.thinkingLevel; + const nextThinkingLevel = selector.thinkingLevel ?? currentThinkingLevel; + + this.#setModelWithProviderSessionReset(candidate); + this.sessionManager.appendModelChange(`${candidate.provider}/${candidate.id}`, "temporary"); + this.settings.getStorage()?.recordModelUsage(`${candidate.provider}/${candidate.id}`); + this.setThinkingLevel(nextThinkingLevel); + if (!this.#activeRetryFallback) { + this.#activeRetryFallback = { + role, + originalSelector: currentSelector, + originalThinkingLevel: currentThinkingLevel, + lastAppliedFallbackThinkingLevel: nextThinkingLevel, + }; + } else { + this.#activeRetryFallback.lastAppliedFallbackThinkingLevel = nextThinkingLevel; + } + await this.#emitSessionEvent({ + type: "retry_fallback_applied", + from: currentSelector, + to: selector.raw, + role, + }); + } + + async #tryRetryModelFallback(currentSelector: string): Promise { + const role = this.#activeRetryFallback?.role ?? this.#resolveRetryFallbackRole(currentSelector); + if (!role) return false; + + for (const selector of this.#findRetryFallbackCandidates(role, currentSelector)) { + if (this.#isRetryFallbackSelectorSuppressed(selector)) continue; + const candidate = this.#modelRegistry.find(selector.provider, selector.id); + if (!candidate) continue; + const apiKey = await this.#modelRegistry.getApiKey(candidate, this.sessionId); + if (!apiKey) continue; + await this.#applyRetryFallbackCandidate(role, selector, currentSelector); + return true; + } + + return false; + } + + async #maybeRestoreRetryFallbackPrimary(): Promise { + if (!this.#activeRetryFallback) return; + if (this.#getRetryFallbackRevertPolicy() !== "cooldown-expiry") return; + + const { + originalSelector: originalSelectorRaw, + originalThinkingLevel, + lastAppliedFallbackThinkingLevel, + } = this.#activeRetryFallback; + const originalSelector = parseRetryFallbackSelector(originalSelectorRaw); + if (!originalSelector) { + this.#clearActiveRetryFallback(); + return; + } + + const currentModel = this.model; + if (!currentModel) return; + const currentSelector = formatRetryFallbackSelector(currentModel, this.thinkingLevel); + if (currentSelector === originalSelector.raw) { + if (!this.#isRetryFallbackSelectorSuppressed(originalSelector)) { + this.#clearActiveRetryFallback(); + } + return; + } + if (this.#isRetryFallbackSelectorSuppressed(originalSelector)) return; + + const primaryModel = this.#modelRegistry.find(originalSelector.provider, originalSelector.id); + if (!primaryModel) return; + const apiKey = await this.#modelRegistry.getApiKey(primaryModel, this.sessionId); + if (!apiKey) return; + + const currentThinkingLevel = this.thinkingLevel; + const thinkingToApply = + currentThinkingLevel === lastAppliedFallbackThinkingLevel ? originalThinkingLevel : currentThinkingLevel; + this.#setModelWithProviderSessionReset(primaryModel); + this.sessionManager.appendModelChange(`${primaryModel.provider}/${primaryModel.id}`, "temporary"); + this.settings.getStorage()?.recordModelUsage(`${primaryModel.provider}/${primaryModel.id}`); + this.setThinkingLevel(thinkingToApply); + this.#clearActiveRetryFallback(); + } + #parseRetryAfterMsFromError(errorMessage: string): number | undefined { const now = Date.now(); const retryAfterMsMatch = /retry-after-ms\s*[:=]\s*(\d+)/i.exec(errorMessage); @@ -4847,12 +5120,13 @@ export class AgentSession { } const errorMessage = message.errorMessage || "Unknown error"; + const parsedRetryAfterMs = this.#parseRetryAfterMsFromError(errorMessage); let delayMs = retrySettings.baseDelayMs * 2 ** (this.#retryAttempt - 1); + let switchedCredential = false; + let switchedModel = false; if (this.model && isUsageLimitError(errorMessage)) { - const retryAfterMs = - this.#parseRetryAfterMsFromError(errorMessage) ?? - calculateRateLimitBackoffMs(parseRateLimitReason(errorMessage)); + const retryAfterMs = parsedRetryAfterMs ?? calculateRateLimitBackoffMs(parseRateLimitReason(errorMessage)); const switched = await this.#modelRegistry.authStorage.markUsageLimitReached( this.model.provider, this.sessionId, @@ -4862,6 +5136,7 @@ export class AgentSession { }, ); if (switched) { + switchedCredential = true; delayMs = 0; } else if (retryAfterMs > delayMs) { // No more accounts to switch to — wait out the backoff @@ -4869,6 +5144,17 @@ export class AgentSession { } } + const currentSelector = this.model ? formatRetryFallbackSelector(this.model, this.thinkingLevel) : undefined; + if (!switchedCredential && currentSelector) { + this.#noteRetryFallbackCooldown(currentSelector, parsedRetryAfterMs, errorMessage); + switchedModel = await this.#tryRetryModelFallback(currentSelector); + if (switchedModel) { + delayMs = 0; + } else if (parsedRetryAfterMs && parsedRetryAfterMs > delayMs) { + delayMs = parsedRetryAfterMs; + } + } + await this.#emitSessionEvent({ type: "auto_retry_start", attempt: this.#retryAttempt, diff --git a/packages/coding-agent/test/agent-session-retry-fallback.test.ts b/packages/coding-agent/test/agent-session-retry-fallback.test.ts new file mode 100644 index 000000000..ec2cf75a2 --- /dev/null +++ b/packages/coding-agent/test/agent-session-retry-fallback.test.ts @@ -0,0 +1,389 @@ +import { afterEach, beforeEach, describe, expect, it } from "bun:test"; +import * as path from "node:path"; +import { Agent } from "@oh-my-pi/pi-agent-core"; +import { type AssistantMessage, Effort, getBundledModel, type Model } from "@oh-my-pi/pi-ai"; +import { AssistantMessageEventStream } from "@oh-my-pi/pi-ai/utils/event-stream"; +import { ModelRegistry } from "@oh-my-pi/pi-coding-agent/config/model-registry"; +import { Settings } from "@oh-my-pi/pi-coding-agent/config/settings"; +import { AgentSession, type AgentSessionEvent } from "@oh-my-pi/pi-coding-agent/session/agent-session"; +import { AuthStorage } from "@oh-my-pi/pi-coding-agent/session/auth-storage"; +import { SessionManager } from "@oh-my-pi/pi-coding-agent/session/session-manager"; +import { TempDir } from "@oh-my-pi/pi-utils"; + +class MockAssistantStream extends AssistantMessageEventStream {} + +function createAssistantMessage( + model: Model, + options: { text?: string; stopReason: "stop" | "error"; errorMessage?: string }, +): AssistantMessage { + return { + role: "assistant", + content: options.text ? [{ type: "text", text: options.text }] : [], + api: model.api, + provider: model.provider, + model: model.id, + usage: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + stopReason: options.stopReason, + errorMessage: options.errorMessage, + timestamp: Date.now(), + }; +} + +async function _waitFor(predicate: () => boolean, timeoutMs = 1000): Promise { + const deadline = Date.now() + timeoutMs; + while (Date.now() < deadline) { + if (predicate()) return; + await Bun.sleep(10); + } + throw new Error("Timed out waiting for condition"); +} + +describe("AgentSession retry fallback", () => { + let tempDir: TempDir; + let authStorage: AuthStorage; + let modelRegistry: ModelRegistry; + let session: AgentSession | undefined; + + beforeEach(async () => { + tempDir = TempDir.createSync("@pi-retry-fallback-"); + authStorage = await AuthStorage.create(path.join(tempDir.path(), "testauth.db")); + authStorage.setRuntimeApiKey("anthropic", "anthropic-test-key"); + authStorage.setRuntimeApiKey("openai", "openai-test-key"); + modelRegistry = new ModelRegistry(authStorage); + }); + + afterEach(async () => { + if (session) { + await session.dispose(); + session = undefined; + } + authStorage.close(); + tempDir.removeSync(); + }); + + it("advances through a role-keyed fallback chain across retries", async () => { + const primaryModel = getBundledModel("anthropic", "claude-sonnet-4-5"); + const firstFallback = getBundledModel("openai", "gpt-4o-mini"); + const secondFallback = getBundledModel("openai", "gpt-4o"); + if (!primaryModel || !firstFallback || !secondFallback) { + throw new Error("Expected bundled test models to exist"); + } + + const requestedModels: string[] = []; + const retryStartEvents: Array> = []; + const retryEndEvents: Array> = []; + const fallbackAppliedEvents: Array> = []; + const fallbackSucceededEvents: Array> = []; + + const agent = new Agent({ + getApiKey: provider => `${provider}-test-key`, + initialState: { + model: primaryModel, + systemPrompt: "Test", + tools: [], + messages: [], + }, + streamFn: model => { + requestedModels.push(`${model.provider}/${model.id}`); + const stream = new MockAssistantStream(); + queueMicrotask(() => { + if (model.provider === primaryModel.provider && model.id === primaryModel.id) { + const message = createAssistantMessage(model, { + stopReason: "error", + errorMessage: "overloaded_error: provider returned error 503", + }); + stream.push({ type: "start", partial: message }); + stream.push({ type: "error", reason: "error", error: message }); + return; + } + if (model.provider === firstFallback.provider && model.id === firstFallback.id) { + const message = createAssistantMessage(model, { + stopReason: "error", + errorMessage: "service unavailable: 503 overloaded", + }); + stream.push({ type: "start", partial: message }); + stream.push({ type: "error", reason: "error", error: message }); + return; + } + if (model.provider === secondFallback.provider && model.id === secondFallback.id) { + const message = createAssistantMessage(model, { + text: "Recovered on second fallback", + stopReason: "stop", + }); + stream.push({ + type: "start", + partial: createAssistantMessage(model, { text: "", stopReason: "stop" }), + }); + stream.push({ type: "done", reason: "stop", message }); + return; + } + throw new Error(`Unexpected model requested during retry fallback test: ${model.provider}/${model.id}`); + }); + return stream; + }, + }); + + const settings = Settings.isolated({ + "compaction.enabled": false, + "retry.baseDelayMs": 5, + "retry.fallbackChains": { + default: [ + `${firstFallback.provider}/${firstFallback.id}`, + `${secondFallback.provider}/${secondFallback.id}`, + ], + }, + }); + settings.setModelRole("default", `${primaryModel.provider}/${primaryModel.id}`); + + session = new AgentSession({ + agent, + sessionManager: SessionManager.inMemory(), + settings, + modelRegistry, + }); + + session.subscribe(event => { + if (event.type === "auto_retry_start") { + retryStartEvents.push(event); + } + if (event.type === "auto_retry_end") { + retryEndEvents.push(event); + } + if (event.type === "retry_fallback_applied") { + fallbackAppliedEvents.push(event); + } + if (event.type === "retry_fallback_succeeded") { + fallbackSucceededEvents.push(event); + } + }); + + await session.prompt("Recover from rate limits"); + await session.waitForIdle(); + + expect(requestedModels).toEqual([ + `${primaryModel.provider}/${primaryModel.id}`, + `${firstFallback.provider}/${firstFallback.id}`, + `${secondFallback.provider}/${secondFallback.id}`, + ]); + expect(session.model?.provider).toBe(secondFallback.provider); + expect(session.model?.id).toBe(secondFallback.id); + expect(retryStartEvents.map(event => event.delayMs)).toEqual([0, 0]); + expect(fallbackAppliedEvents).toEqual([ + { + type: "retry_fallback_applied", + from: `${primaryModel.provider}/${primaryModel.id}`, + to: `${firstFallback.provider}/${firstFallback.id}`, + role: "default", + }, + { + type: "retry_fallback_applied", + from: `${firstFallback.provider}/${firstFallback.id}`, + to: `${secondFallback.provider}/${secondFallback.id}`, + role: "default", + }, + ]); + expect(retryEndEvents).toHaveLength(1); + expect(retryEndEvents[0]).toMatchObject({ success: true, attempt: 2 }); + expect(fallbackSucceededEvents).toEqual([ + { + type: "retry_fallback_succeeded", + model: `${secondFallback.provider}/${secondFallback.id}`, + role: "default", + }, + ]); + }); + + it("suppresses cooled selectors and lazily reverts to the role primary after cooldown expiry", async () => { + const primaryModel = getBundledModel("anthropic", "claude-sonnet-4-5"); + const fallbackModel = getBundledModel("openai", "gpt-4o-mini"); + if (!primaryModel || !fallbackModel) { + throw new Error("Expected bundled test models to exist"); + } + + const requestedModels: string[] = []; + let primaryAttempts = 0; + + const agent = new Agent({ + getApiKey: provider => `${provider}-test-key`, + initialState: { + model: primaryModel, + systemPrompt: "Test", + tools: [], + messages: [], + }, + streamFn: model => { + requestedModels.push(`${model.provider}/${model.id}`); + const stream = new MockAssistantStream(); + queueMicrotask(() => { + if (model.provider === primaryModel.provider && model.id === primaryModel.id && primaryAttempts === 0) { + primaryAttempts += 1; + const message = createAssistantMessage(model, { + stopReason: "error", + errorMessage: "rate limit exceeded retry-after-ms=200", + }); + stream.push({ type: "start", partial: message }); + stream.push({ type: "error", reason: "error", error: message }); + return; + } + const message = createAssistantMessage(model, { + text: `ok:${model.provider}/${model.id}`, + stopReason: "stop", + }); + stream.push({ type: "start", partial: createAssistantMessage(model, { text: "", stopReason: "stop" }) }); + stream.push({ type: "done", reason: "stop", message }); + }); + return stream; + }, + }); + + const settings = Settings.isolated({ + "compaction.enabled": false, + "retry.baseDelayMs": 5, + "retry.fallbackChains": { + default: [`${fallbackModel.provider}/${fallbackModel.id}`], + }, + "retry.fallbackRevertPolicy": "cooldown-expiry", + }); + settings.setModelRole("default", `${primaryModel.provider}/${primaryModel.id}`); + + session = new AgentSession({ + agent, + sessionManager: SessionManager.inMemory(), + settings, + modelRegistry, + }); + + await session.prompt("First prompt triggers fallback"); + await session.waitForIdle(); + expect(requestedModels).toEqual([ + `${primaryModel.provider}/${primaryModel.id}`, + `${fallbackModel.provider}/${fallbackModel.id}`, + ]); + expect(session.model?.provider).toBe(fallbackModel.provider); + expect(session.model?.id).toBe(fallbackModel.id); + + await session.prompt("Immediate second prompt should stay on fallback"); + await session.waitForIdle(); + expect(requestedModels).toEqual([ + `${primaryModel.provider}/${primaryModel.id}`, + `${fallbackModel.provider}/${fallbackModel.id}`, + `${fallbackModel.provider}/${fallbackModel.id}`, + ]); + expect(session.model?.provider).toBe(fallbackModel.provider); + expect(session.model?.id).toBe(fallbackModel.id); + + await Bun.sleep(240); + await session.prompt("Third prompt should lazily revert to primary"); + await session.waitForIdle(); + expect(requestedModels).toEqual([ + `${primaryModel.provider}/${primaryModel.id}`, + `${fallbackModel.provider}/${fallbackModel.id}`, + `${fallbackModel.provider}/${fallbackModel.id}`, + `${primaryModel.provider}/${primaryModel.id}`, + ]); + expect(session.model?.provider).toBe(primaryModel.provider); + expect(session.model?.id).toBe(primaryModel.id); + }); + + it("preserves thinking on bare fallback selectors and does not overwrite user thinking on restore", async () => { + const primaryModel = getBundledModel("anthropic", "claude-sonnet-4-5"); + const fallbackModel = getBundledModel("openai", "gpt-4o-mini"); + if (!primaryModel || !fallbackModel) { + throw new Error("Expected bundled test models to exist"); + } + + const requestedModels: string[] = []; + let primaryAttempts = 0; + + const agent = new Agent({ + getApiKey: provider => `${provider}-test-key`, + initialState: { + model: primaryModel, + systemPrompt: "Test", + tools: [], + messages: [], + }, + streamFn: model => { + requestedModels.push(`${model.provider}/${model.id}`); + const stream = new MockAssistantStream(); + queueMicrotask(() => { + if (model.provider === primaryModel.provider && model.id === primaryModel.id && primaryAttempts === 0) { + primaryAttempts += 1; + const message = createAssistantMessage(model, { + stopReason: "error", + errorMessage: "rate limit exceeded retry-after-ms=200", + }); + stream.push({ type: "start", partial: message }); + stream.push({ type: "error", reason: "error", error: message }); + return; + } + const message = createAssistantMessage(model, { + text: `ok:${model.provider}/${model.id}`, + stopReason: "stop", + }); + stream.push({ type: "start", partial: createAssistantMessage(model, { text: "", stopReason: "stop" }) }); + stream.push({ type: "done", reason: "stop", message }); + }); + return stream; + }, + }); + + const settings = Settings.isolated({ + "compaction.enabled": false, + "retry.baseDelayMs": 5, + "retry.fallbackChains": { + default: [`${fallbackModel.provider}/${fallbackModel.id}`], + }, + "retry.fallbackRevertPolicy": "cooldown-expiry", + }); + settings.setModelRole("default", `${primaryModel.provider}/${primaryModel.id}:high`); + + session = new AgentSession({ + agent, + sessionManager: SessionManager.inMemory(), + settings, + modelRegistry, + thinkingLevel: Effort.High, + }); + + await session.prompt("First prompt triggers bare-selector fallback"); + await session.waitForIdle(); + expect(requestedModels).toEqual([ + `${primaryModel.provider}/${primaryModel.id}`, + `${fallbackModel.provider}/${fallbackModel.id}`, + ]); + expect(session.model?.provider).toBe(fallbackModel.provider); + expect(session.model?.id).toBe(fallbackModel.id); + expect(session.thinkingLevel).toBeUndefined(); + + session.setThinkingLevel(Effort.Low); + await Bun.sleep(240); + await session.prompt("Second prompt should restore model but preserve user thinking change"); + await session.waitForIdle(); + expect(requestedModels).toEqual([ + `${primaryModel.provider}/${primaryModel.id}`, + `${fallbackModel.provider}/${fallbackModel.id}`, + `${primaryModel.provider}/${primaryModel.id}`, + ]); + expect(session.model?.provider).toBe(primaryModel.provider); + expect(session.model?.id).toBe(primaryModel.id); + expect(session.thinkingLevel).toBeUndefined(); + }); + + it("normalizes suppression by base selector and clears it on model refresh", async () => { + const future = Date.now() + 60_000; + modelRegistry.suppressSelector("openai/gpt-4o:high", future); + expect(modelRegistry.isSelectorSuppressed("openai/gpt-4o")).toBe(true); + expect(modelRegistry.isSelectorSuppressed("openai/gpt-4o:low")).toBe(true); + + await modelRegistry.refresh("offline"); + expect(modelRegistry.isSelectorSuppressed("openai/gpt-4o")).toBe(false); + }); +});