diff --git a/src/hooks/todo-continuation-enforcer/handler.ts b/src/hooks/todo-continuation-enforcer/handler.ts index 7136dda44..27096056f 100644 --- a/src/hooks/todo-continuation-enforcer/handler.ts +++ b/src/hooks/todo-continuation-enforcer/handler.ts @@ -13,6 +13,45 @@ import { handleSessionIdle } from "./idle-event" import { handleNonIdleEvent } from "./non-idle-events" import { isTokenLimitError } from "./token-limit-detection" +function asRecord(value: unknown): Record | undefined { + return typeof value === "object" && value !== null ? value as Record : undefined +} + +function getStringField(record: Record | undefined, key: string): string | undefined { + const value = record?.[key] + return typeof value === "string" && value.length > 0 ? value : undefined +} + +function extractSessionErrorInfo(error: unknown): { name?: string; message?: string } | undefined { + if (!error) return undefined + if (typeof error === "string") return { message: error } + if (error instanceof Error) return { name: error.name, message: error.message } + + const root = asRecord(error) + if (!root) return { message: String(error) } + + const data = asRecord(root.data) + const nestedError = asRecord(root.error) + const dataError = asRecord(data?.error) + + const name = getStringField(root, "name") + ?? getStringField(data, "name") + ?? getStringField(nestedError, "name") + ?? getStringField(dataError, "name") + + const messageParts = [ + getStringField(root, "message"), + getStringField(data, "message"), + getStringField(nestedError, "message"), + getStringField(dataError, "message"), + getStringField(root, "code"), + getStringField(nestedError, "code"), + getStringField(dataError, "code"), + ].filter((message): message is string => typeof message === "string") + + return { name, message: messageParts.join(" ") || undefined } +} + export function createTodoContinuationHandler(args: { ctx: PluginInput sessionStateStore: SessionStateStore @@ -35,7 +74,8 @@ export function createTodoContinuationHandler(args: { const sessionID = props?.sessionID as string | undefined if (!sessionID) return - const error = props?.error as { name?: string; message?: string } | undefined + const error = extractSessionErrorInfo(props?.error) + let shouldCancelCountdown = false if (error?.name === "MessageAbortedError" || error?.name === "AbortError") { const state = sessionStateStore.getState(sessionID) state.wasCancelled = true @@ -45,14 +85,18 @@ export function createTodoContinuationHandler(args: { state.awaitingPostInjectionProgressCheck = false state.stagnationCount = 0 state.consecutiveFailures = 0 + shouldCancelCountdown = true log(`[${HOOK_NAME}] Abort detected via session.error`, { sessionID, errorName: error.name }) } else if (isTokenLimitError(error)) { const state = sessionStateStore.getState(sessionID) state.tokenLimitDetected = true + shouldCancelCountdown = true log(`[${HOOK_NAME}] Token limit error detected via session.error`, { sessionID, errorName: error?.name, errorMessage: error?.message }) } - sessionStateStore.cancelCountdown(sessionID) + if (shouldCancelCountdown) { + sessionStateStore.cancelCountdown(sessionID) + } log(`[${HOOK_NAME}] session.error`, { sessionID }) return } diff --git a/src/hooks/todo-continuation-enforcer/opencode-overload-continuation.test.ts b/src/hooks/todo-continuation-enforcer/opencode-overload-continuation.test.ts new file mode 100644 index 000000000..45686729e --- /dev/null +++ b/src/hooks/todo-continuation-enforcer/opencode-overload-continuation.test.ts @@ -0,0 +1,89 @@ +import { describe, expect, test } from "bun:test" + +import { _resetForTesting, setMainSession } from "../../features/claude-code-session-state" +import { createTodoContinuationEnforcer } from "." + +type PromptCall = { + sessionID: string + text: string +} + +type PromptInput = { + path: { id: string } + body: { parts: Array<{ text: string }> } +} + +function wait(ms: number): Promise { + return new Promise((resolve) => setTimeout(resolve, ms)) +} + +function createPluginInput(promptCalls: PromptCall[]): Parameters[0] { + return { + directory: "/tmp/opencode-overload-continuation-test", + client: { + session: { + todo: async () => ({ + data: [ + { id: "1", content: "Keep working", status: "pending", priority: "high" }, + ], + }), + messages: async () => ({ data: [] }), + promptAsync: async (input: PromptInput) => { + promptCalls.push({ + sessionID: input.path.id, + text: input.body.parts[0]?.text ?? "", + }) + return {} + }, + }, + tui: { + showToast: async () => ({}), + }, + }, + } as Parameters[0] +} + +describe("todo-continuation-enforcer OpenCode overload errors", () => { + test( + "#given countdown is armed #when OpenCode reports server_is_overloaded #then continuation still injects", + async () => { + // given + const sessionID = "main-opencode-overload" + const promptCalls: PromptCall[] = [] + _resetForTesting() + setMainSession(sessionID) + const hook = createTodoContinuationEnforcer(createPluginInput(promptCalls)) + + await hook.handler({ + event: { type: "session.idle", properties: { sessionID } }, + }) + + // when + await hook.handler({ + event: { + type: "session.error", + properties: { + sessionID, + error: { + type: "error", + sequence_number: 2, + error: { + type: "service_unavailable_error", + code: "server_is_overloaded", + message: "Our servers are currently overloaded. Please try again later.", + param: null, + }, + }, + }, + }, + }) + await wait(2500) + + // then + expect(promptCalls).toHaveLength(1) + expect(promptCalls[0]?.sessionID).toBe(sessionID) + expect(promptCalls[0]?.text).toContain("TODO CONTINUATION") + }, + { timeout: 10000 }, + ) +})