diff --git a/cli/src/claude/claudeRemote.seam.test.ts b/cli/src/claude/claudeRemote.seam.test.ts new file mode 100644 index 00000000..bc4fda15 --- /dev/null +++ b/cli/src/claude/claudeRemote.seam.test.ts @@ -0,0 +1,115 @@ +import { EventEmitter } from 'node:events' +import { PassThrough } from 'node:stream' +import { afterEach, describe, expect, it, vi } from 'vitest' +import type { SDKMessage } from '@/claude/sdk/types' + +const spawnMock = vi.fn() +const killProcessMock = vi.fn(async (child: any) => { + child.killed = true + child.stdout.end() + child.emit('close', 0) + return true +}) + +vi.mock('node:child_process', () => ({ + ...require('node:child_process'), + spawn: spawnMock +})) + +vi.mock('@/claude/utils/claudeCheckSession', () => ({ + claudeCheckSession: () => true +})) + +vi.mock('@/modules/watcher/awaitFileExist', () => ({ + awaitFileExist: async () => true +})) + +vi.mock('@/utils/process', () => ({ + isProcessAlive: () => false, + isWindows: () => false, + killProcess: async () => true, + killProcessByChildProcess: killProcessMock +})) + +vi.mock('@/utils/bunRuntime', () => ({ + withBunRuntimeEnv: (env: NodeJS.ProcessEnv) => env +})) + +function createFakeChild() { + const child = new EventEmitter() as EventEmitter & { + stdin: PassThrough + stdout: PassThrough + stderr: PassThrough + killed: boolean + } + + child.stdin = new PassThrough() + child.stdout = new PassThrough() + child.stderr = new PassThrough() + child.killed = false + return child +} + +afterEach(() => { + vi.clearAllMocks() + delete process.env.HAPI_CLAUDE_PATH +}) + +describe('claudeRemote/query real seam', () => { + it('propagates scheduled nextMessage failures through real query prompt plumbing', async () => { + const child = createFakeChild() + spawnMock.mockReturnValueOnce(child) + process.env.HAPI_CLAUDE_PATH = 'claude' + const { claudeRemote } = await import('./claudeRemote') + + const received: SDKMessage[] = [] + let nextCallCount = 0 + + const runPromise = claudeRemote({ + sessionId: 'session-1', + path: process.cwd(), + mcpServers: {}, + claudeEnvVars: {}, + claudeArgs: [], + allowedTools: [], + hookSettingsPath: '/tmp/hook.json', + canCallTool: async () => ({ behavior: 'allow', updatedInput: {} }), + nextMessage: async () => { + nextCallCount += 1 + if (nextCallCount === 1) { + return { message: 'A', mode: { permissionMode: 'default' } } + } + throw new Error('next message failed') + }, + onReady: () => {}, + isAborted: () => false, + onSessionFound: () => {}, + onMessage: (message) => { + received.push(message) + }, + onCompletionEvent: () => {}, + onSessionReset: () => {} + }) + + child.stdout.write(JSON.stringify({ + type: 'assistant', + message: { + role: 'assistant', + content: [{ type: 'text', text: 'A_1' }] + } + }) + '\n') + child.stdout.write(JSON.stringify({ + type: 'result', + subtype: 'success', + num_turns: 1, + total_cost_usd: 0, + duration_ms: 1, + duration_api_ms: 1, + is_error: false, + session_id: 's-1' + }) + '\n') + + await expect(runPromise).rejects.toThrow('next message failed') + expect(received.map((message) => message.type)).toEqual(['assistant', 'result']) + }) +}) diff --git a/cli/src/claude/claudeRemote.test.ts b/cli/src/claude/claudeRemote.test.ts new file mode 100644 index 00000000..dd0a11cc --- /dev/null +++ b/cli/src/claude/claudeRemote.test.ts @@ -0,0 +1,287 @@ +import { describe, it, expect, vi } from 'vitest'; +import * as claudeSdk from '@/claude/sdk'; +import type { SDKMessage } from '@/claude/sdk/types'; + +vi.mock('@/claude/utils/claudeCheckSession', () => ({ + claudeCheckSession: () => true +})); + +vi.mock('@/modules/watcher/awaitFileExist', () => ({ + awaitFileExist: async () => true +})); + +vi.mock('@/claude/sdk/utils', () => ({ + getDefaultClaudeCodePath: () => '/usr/bin/claude' +})); + +const queryMock = vi.fn(); + +function deferred(): { + promise: Promise; + resolve: (value: T) => void; + reject: (reason?: unknown) => void; +} { + let resolve!: (value: T) => void; + let reject!: (reason?: unknown) => void; + const promise = new Promise((res, rej) => { + resolve = res; + reject = rej; + }); + return { promise, resolve, reject }; +} + +function createAsyncStream(messages: SDKMessage[]): AsyncIterable { + return { + async *[Symbol.asyncIterator]() { + for (const message of messages) { + await Promise.resolve(); + yield message; + } + } + }; +} + +function createQueryThatMirrorsPromptErrors(messages: SDKMessage[]) { + return ({ prompt }: { prompt: AsyncIterable }) => ({ + async *[Symbol.asyncIterator]() { + const promptIterator = prompt[Symbol.asyncIterator](); + + await promptIterator.next(); + + for (const message of messages) { + await Promise.resolve(); + yield message; + } + + await promptIterator.next(); + } + }); +} + +async function waitFor(condition: () => boolean, timeoutMs = 300, intervalMs = 10): Promise { + const startedAt = Date.now(); + while (!condition()) { + if (Date.now() - startedAt > timeoutMs) { + throw new Error('Timed out waiting for condition'); + } + await new Promise((resolve) => setTimeout(resolve, intervalMs)); + } +} + +describe('claudeRemote async message handling', () => { + it('continues consuming assistant messages even when next user message is pending', async () => { + const querySpy = vi.spyOn(claudeSdk, 'query').mockImplementation(queryMock as typeof claudeSdk.query); + const { claudeRemote } = await import('./claudeRemote'); + const pendingNext = deferred<{ message: string; mode: { permissionMode: 'default' } } | null>(); + const received: SDKMessage[] = []; + + const sdkMessages: SDKMessage[] = [ + { + type: 'assistant', + message: { + role: 'assistant', + content: [{ type: 'text', text: 'A_1' }] + } + } as unknown as SDKMessage, + { + type: 'result', + subtype: 'success', + num_turns: 1, + total_cost_usd: 0, + duration_ms: 1, + duration_api_ms: 1, + is_error: false, + session_id: 's-1' + } as unknown as SDKMessage, + { + type: 'assistant', + message: { + role: 'assistant', + content: [{ type: 'text', text: 'A_2' }] + } + } as unknown as SDKMessage + ]; + + queryMock.mockReturnValueOnce(createAsyncStream(sdkMessages)); + + let nextCallCount = 0; + const runPromise = claudeRemote({ + sessionId: 'session-1', + path: process.cwd(), + mcpServers: {}, + claudeEnvVars: {}, + claudeArgs: [], + allowedTools: [], + hookSettingsPath: '/tmp/hook.json', + canCallTool: async () => ({ behavior: 'allow', updatedInput: {} }), + nextMessage: async () => { + nextCallCount += 1; + if (nextCallCount === 1) { + return { message: 'A', mode: { permissionMode: 'default' } }; + } + return await pendingNext.promise; + }, + onReady: () => {}, + isAborted: () => false, + onSessionFound: () => {}, + onMessage: (message) => { + received.push(message); + }, + onCompletionEvent: () => {}, + onSessionReset: () => {} + }); + + await waitFor(() => received.length >= 3); + expect(received.map((m) => m.type)).toEqual(['assistant', 'result', 'assistant']); + + try { + pendingNext.resolve(null); + await runPromise; + } finally { + queryMock.mockReset(); + querySpy.mockRestore(); + } + }); + + it('handles rejected next user message fetch without unhandled rejection', async () => { + const querySpy = vi.spyOn(claudeSdk, 'query').mockImplementation(queryMock as typeof claudeSdk.query); + const { claudeRemote } = await import('./claudeRemote'); + const received: SDKMessage[] = []; + const unhandled: unknown[] = []; + const onUnhandled = (reason: unknown) => { + unhandled.push(reason); + }; + process.on('unhandledRejection', onUnhandled); + + const sdkMessages: SDKMessage[] = [ + { + type: 'assistant', + message: { + role: 'assistant', + content: [{ type: 'text', text: 'A_1' }] + } + } as unknown as SDKMessage, + { + type: 'result', + subtype: 'success', + num_turns: 1, + total_cost_usd: 0, + duration_ms: 1, + duration_api_ms: 1, + is_error: false, + session_id: 's-1' + } as unknown as SDKMessage + ]; + + queryMock.mockImplementationOnce(createQueryThatMirrorsPromptErrors(sdkMessages)); + + let nextCallCount = 0; + const runPromise = claudeRemote({ + sessionId: 'session-1', + path: process.cwd(), + mcpServers: {}, + claudeEnvVars: {}, + claudeArgs: [], + allowedTools: [], + hookSettingsPath: '/tmp/hook.json', + canCallTool: async () => ({ behavior: 'allow', updatedInput: {} }), + nextMessage: async () => { + nextCallCount += 1; + if (nextCallCount === 1) { + return { message: 'A', mode: { permissionMode: 'default' } }; + } + throw new Error('next message failed'); + }, + onReady: () => {}, + isAborted: () => false, + onSessionFound: () => {}, + onMessage: (message) => { + received.push(message); + }, + onCompletionEvent: () => {}, + onSessionReset: () => {} + }); + + try { + await expect(runPromise).rejects.toThrow('next message failed'); + await new Promise((resolve) => setTimeout(resolve, 0)); + expect(received.map((m) => m.type)).toEqual(['assistant', 'result']); + expect(unhandled).toEqual([]); + } finally { + queryMock.mockReset(); + querySpy.mockRestore(); + process.off('unhandledRejection', onUnhandled); + } + }); + + it('treats AbortError from scheduled next user message fetch as graceful shutdown', async () => { + const querySpy = vi.spyOn(claudeSdk, 'query').mockImplementation(queryMock as typeof claudeSdk.query); + const { claudeRemote } = await import('./claudeRemote'); + const received: SDKMessage[] = []; + const unhandled: unknown[] = []; + const onUnhandled = (reason: unknown) => { + unhandled.push(reason); + }; + process.on('unhandledRejection', onUnhandled); + + const sdkMessages: SDKMessage[] = [ + { + type: 'assistant', + message: { + role: 'assistant', + content: [{ type: 'text', text: 'A_1' }] + } + } as unknown as SDKMessage, + { + type: 'result', + subtype: 'success', + num_turns: 1, + total_cost_usd: 0, + duration_ms: 1, + duration_api_ms: 1, + is_error: false, + session_id: 's-1' + } as unknown as SDKMessage + ]; + + queryMock.mockReturnValueOnce(createAsyncStream(sdkMessages)); + + let nextCallCount = 0; + const runPromise = claudeRemote({ + sessionId: 'session-1', + path: process.cwd(), + mcpServers: {}, + claudeEnvVars: {}, + claudeArgs: [], + allowedTools: [], + hookSettingsPath: '/tmp/hook.json', + canCallTool: async () => ({ behavior: 'allow', updatedInput: {} }), + nextMessage: async () => { + nextCallCount += 1; + if (nextCallCount === 1) { + return { message: 'A', mode: { permissionMode: 'default' } }; + } + throw new claudeSdk.AbortError('aborted'); + }, + onReady: () => {}, + isAborted: () => false, + onSessionFound: () => {}, + onMessage: (message) => { + received.push(message); + }, + onCompletionEvent: () => {}, + onSessionReset: () => {} + }); + + try { + await runPromise; + await new Promise((resolve) => setTimeout(resolve, 0)); + expect(received.map((m) => m.type)).toEqual(['assistant', 'result']); + expect(unhandled).toEqual([]); + } finally { + queryMock.mockReset(); + querySpy.mockRestore(); + process.off('unhandledRejection', onUnhandled); + } + }); +}); diff --git a/cli/src/claude/claudeRemote.ts b/cli/src/claude/claudeRemote.ts index 9ca08619..038b037f 100644 --- a/cli/src/claude/claudeRemote.ts +++ b/cli/src/claude/claudeRemote.ts @@ -37,6 +37,7 @@ export async function claudeRemote(opts: { onCompletionEvent?: (message: string) => void, onSessionReset?: () => void }) { + const debugPrefix = '[claudeRemote][async-debug]'; // Check if session is valid let startFrom = opts.sessionId; @@ -90,8 +91,10 @@ export async function claudeRemote(opts: { throw e; } if (!initial) { // No initial message - exit + logger.debug(`${debugPrefix} initial nextMessage returned null; exiting`); return; } + logger.debug(`${debugPrefix} initial message acquired`); // Handle special commands const specialCommand = parseSpecialCommand(initial.message); @@ -166,11 +169,68 @@ export async function claudeRemote(opts: { options: sdkOptions, }); + let nextMessageFetchInFlight = false; + let inputEnded = false; + let nextMessageFetchSeq = 0; + let streamMessageSeq = 0; + let resultSeq = 0; + + const scheduleNextMessage = () => { + if (nextMessageFetchInFlight || inputEnded) { + logger.debug( + `${debugPrefix} scheduleNextMessage skipped ` + + `(inFlight=${nextMessageFetchInFlight}, inputEnded=${inputEnded})` + ); + return; + } + + const fetchId = ++nextMessageFetchSeq; + const startedAt = Date.now(); + nextMessageFetchInFlight = true; + logger.debug(`${debugPrefix} scheduleNextMessage start fetchId=${fetchId}`); + void (async () => { + try { + const next = await opts.nextMessage(); + if (!next) { + inputEnded = true; + messages.end(); + logger.debug( + `${debugPrefix} nextMessage resolved null fetchId=${fetchId} elapsedMs=${Date.now() - startedAt}; input ended` + ); + return; + } + mode = next.mode; + messages.push({ type: 'user', message: { role: 'user', content: next.message } }); + logger.debug( + `${debugPrefix} nextMessage resolved fetchId=${fetchId} elapsedMs=${Date.now() - startedAt} ` + + `messageLength=${next.message.length} permissionMode=${next.mode.permissionMode}` + ); + } catch (e) { + inputEnded = true; + if (e instanceof AbortError) { + messages.end(); + logger.debug(`${debugPrefix} nextMessage aborted fetchId=${fetchId}`); + return; + } + messages.setError(e instanceof Error ? e : new Error(String(e))); + logger.debug(`${debugPrefix} nextMessage error fetchId=${fetchId}`, e); + } finally { + nextMessageFetchInFlight = false; + logger.debug(`${debugPrefix} scheduleNextMessage done fetchId=${fetchId}`); + } + })(); + }; + updateThinking(true); try { logger.debug(`[claudeRemote] Starting to iterate over response`); for await (const message of response) { + streamMessageSeq += 1; + logger.debug( + `${debugPrefix} stream message #${streamMessageSeq} type=${message.type} ` + + `subtype=${'subtype' in message ? String((message as any).subtype) : 'n/a'}` + ); logger.debugLargeJson(`[claudeRemote] Message ${message.type}`, message); // Handle messages @@ -196,8 +256,12 @@ export async function claudeRemote(opts: { // Handle result messages if (message.type === 'result') { + resultSeq += 1; updateThinking(false); - logger.debug('[claudeRemote] Result received, exiting claudeRemote'); + logger.debug( + `${debugPrefix} result #${resultSeq} received; scheduling next user message ` + + `(nextInFlight=${nextMessageFetchInFlight}, inputEnded=${inputEnded})` + ); // Send completion messages if (isCompactCommand) { @@ -210,15 +274,12 @@ export async function claudeRemote(opts: { // Send ready event opts.onReady(); + logger.debug(`${debugPrefix} onReady emitted for result #${resultSeq}`); - // Push next message - const next = await opts.nextMessage(); - if (!next) { - messages.end(); - return; - } - mode = next.mode; - messages.push({ type: 'user', message: { role: 'user', content: next.message } }); + // Pull next user message without blocking response stream processing. + // Claude may emit autonomous async messages (e.g. scheduled tasks) after a result, + // and we must keep consuming those messages immediately. + scheduleNextMessage(); } // Handle tool result @@ -228,20 +289,27 @@ export async function claudeRemote(opts: { for (let c of msg.message.content) { if (c.type === 'tool_result' && c.tool_use_id && opts.isAborted(c.tool_use_id)) { logger.debug('[claudeRemote] Tool aborted, exiting claudeRemote'); + logger.debug(`${debugPrefix} tool aborted via tool_result; exiting stream loop`); return; } } } } } + logger.debug(`${debugPrefix} response stream exhausted`); } catch (e) { if (e instanceof AbortError) { logger.debug(`[claudeRemote] Aborted`); // Ignore } else { + logger.debug(`${debugPrefix} response stream error`, e); throw e; } } finally { + logger.debug( + `${debugPrefix} finally ` + + `(streamMessages=${streamMessageSeq}, results=${resultSeq}, nextFetches=${nextMessageFetchSeq}, inputEnded=${inputEnded})` + ); updateThinking(false); } } diff --git a/cli/src/claude/claudeRemoteLauncher.ts b/cli/src/claude/claudeRemoteLauncher.ts index 477336cd..c5f4a327 100644 --- a/cli/src/claude/claudeRemoteLauncher.ts +++ b/cli/src/claude/claudeRemoteLauncher.ts @@ -348,8 +348,15 @@ class ClaudeRemoteLauncher extends RemoteLauncherBase { session.clearSessionId(); }, onReady: () => { + logger.debug( + `[claudeRemoteLauncher][async-debug] onReady callback ` + + `(hasPending=${Boolean(pending)}, queueSize=${session.queue.size()})` + ); if (!pending && session.queue.size() === 0) { session.client.sendSessionEvent({ type: 'ready' }); + logger.debug('[claudeRemoteLauncher][async-debug] ready event sent to hub'); + } else { + logger.debug('[claudeRemoteLauncher][async-debug] ready event suppressed (pending input exists)'); } }, signal: controller.signal, diff --git a/cli/src/claude/sdk/query.test.ts b/cli/src/claude/sdk/query.test.ts new file mode 100644 index 00000000..72860249 --- /dev/null +++ b/cli/src/claude/sdk/query.test.ts @@ -0,0 +1,100 @@ +import { EventEmitter } from 'node:events' +import { PassThrough } from 'node:stream' +import { afterEach, describe, expect, it, vi } from 'vitest' + +const spawnMock = vi.fn() +const killProcessMock = vi.fn(async (child: any) => { + child.killed = true + child.stdout.end() + child.emit('close', 0) +}) + +vi.mock('node:child_process', () => ({ + ...require('node:child_process'), + spawn: spawnMock +})) + +vi.mock('@/utils/process', () => ({ + isProcessAlive: () => false, + isWindows: () => false, + killProcess: async () => true, + killProcessByChildProcess: killProcessMock +})) + +vi.mock('@/utils/bunRuntime', () => ({ + withBunRuntimeEnv: (env: NodeJS.ProcessEnv) => env +})) + +vi.mock('../utils/mcpConfig', () => ({ + appendMcpConfigArg: () => null +})) + +function createFakeChild() { + const child = new EventEmitter() as EventEmitter & { + stdin: PassThrough + stdout: PassThrough + stderr: PassThrough + killed: boolean + } + + child.stdin = new PassThrough() + child.stdout = new PassThrough() + child.stderr = new PassThrough() + child.killed = false + return child +} + +afterEach(() => { + vi.clearAllMocks() + delete process.env.HAPI_CLAUDE_PATH +}) + +describe('Query', () => { + it('preserves externally set errors even if the process exits cleanly', async () => { + const { Query } = await import('./query') + const stdout = new PassThrough() + const query = new Query(null, stdout, Promise.resolve()) + + query.setError(new Error('prompt failed')) + stdout.end() + + await expect(query.next()).rejects.toThrow('prompt failed') + }) + + it('propagates prompt stream failures through query()', async () => { + const child = createFakeChild() + spawnMock.mockReturnValueOnce(child) + process.env.HAPI_CLAUDE_PATH = 'claude' + + const { query } = await import('./query') + const prompt = { + async *[Symbol.asyncIterator]() { + yield { type: 'user', message: { role: 'user', content: 'hello' } } + throw new Error('prompt failed') + } + } + + const result = query({ prompt }) + + await expect(result.next()).rejects.toThrow('prompt failed') + }) + + it('fails fast after cleanup timeout when prompt cleanup hangs', async () => { + const child = createFakeChild() + spawnMock.mockReturnValueOnce(child) + killProcessMock.mockReturnValueOnce(new Promise(() => {})) + process.env.HAPI_CLAUDE_PATH = 'claude' + + const { query } = await import('./query') + const prompt = { + async *[Symbol.asyncIterator]() { + yield { type: 'user', message: { role: 'user', content: 'hello' } } + throw new Error('prompt failed') + } + } + + const result = query({ prompt, options: { promptFailureCleanupTimeoutMs: 10 } }) + + await expect(result.next()).rejects.toThrow('prompt failed') + }) +}) diff --git a/cli/src/claude/sdk/query.ts b/cli/src/claude/sdk/query.ts index 2f68fc13..246fb32a 100644 --- a/cli/src/claude/sdk/query.ts +++ b/cli/src/claude/sdk/query.ts @@ -30,6 +30,8 @@ import type { Writable } from 'node:stream' import { logger } from '@/ui/logger' import { appendMcpConfigArg } from '../utils/mcpConfig' +const DEFAULT_PROMPT_FAILURE_CLEANUP_TIMEOUT_MS = 3_000 + /** * Query class manages Claude Code process interaction */ @@ -39,6 +41,7 @@ export class Query implements AsyncIterableIterator { private sdkMessages: AsyncIterableIterator private inputStream = new Stream() private canCallTool?: CanCallToolCallback + private promptFailure: Error | null = null constructor( private childStdin: Writable | null, @@ -58,6 +61,19 @@ export class Query implements AsyncIterableIterator { this.inputStream.error(error) } + registerPromptFailure(error: Error): boolean { + if (this.promptFailure) { + return false + } + this.promptFailure = error + this.cleanupControllers() + return true + } + + getPromptFailure(): Error | null { + return this.promptFailure + } + /** * AsyncIterableIterator implementation */ @@ -92,10 +108,18 @@ export class Query implements AsyncIterableIterator { try { for await (const line of rl) { + if (this.promptFailure) { + break + } + if (line.trim()) { try { const message = JSON.parse(line) as SDKMessage | SDKControlResponse + if (this.promptFailure) { + break + } + if (message.type === 'control_response') { const controlResponse = message as SDKControlResponse const handler = this.pendingControlResponses.get(controlResponse.response.request_id) @@ -124,7 +148,7 @@ export class Query implements AsyncIterableIterator { } finally { // Only call done() on clean exit - calling done() after error() // would mask the error since Stream.next() checks isDone before hasError - if (!hadError) { + if (!hadError && !this.inputStream.hasTerminalError) { this.inputStream.done() } this.cleanupControllers() @@ -193,6 +217,9 @@ export class Query implements AsyncIterableIterator { try { const response = await this.processControlRequest(request, controller.signal) + if (this.promptFailure || controller.signal.aborted || !this.childStdin?.writable) { + return + } const controlResponse: CanUseToolControlResponse = { type: 'control_response', response: { @@ -203,6 +230,9 @@ export class Query implements AsyncIterableIterator { } this.childStdin.write(JSON.stringify(controlResponse) + '\n') } catch (error) { + if (this.promptFailure || controller.signal.aborted || !this.childStdin?.writable) { + return + } const controlErrorResponse: CanUseToolControlResponse = { type: 'control_response', response: { @@ -284,7 +314,8 @@ export function query(config: { fallbackModel, settingsPath, strictMcpConfig, - canCallTool + canCallTool, + promptFailureCleanupTimeoutMs = DEFAULT_PROMPT_FAILURE_CLEANUP_TIMEOUT_MS } = {} } = config @@ -362,12 +393,19 @@ export function query(config: { windowsHide: process.platform === 'win32' }) as ChildProcessWithoutNullStreams + // Handle process exit + let resolveExit: () => void + let rejectExit: (error: Error) => void + const processExitPromise = new Promise((resolve, reject) => { + resolveExit = resolve + rejectExit = reject + }) + // Handle stdin let childStdin: Writable | null = null if (typeof prompt === 'string') { child.stdin.end() } else { - streamToStdin(prompt, child.stdin, config.options?.abort) childStdin = child.stdin } @@ -379,30 +417,54 @@ export function query(config: { } // Setup cleanup - const cleanup = () => { - if (!child.killed) { - void killProcessByChildProcess(child) + let cleanupPromise: Promise | null = null + const cleanup = (): Promise => { + if (cleanupPromise) { + return cleanupPromise } + cleanupPromise = (async () => { + await killProcessByChildProcess(child) + child.stdin.destroy() + child.stdout.destroy() + child.stderr.destroy() + })() + return cleanupPromise } - config.options?.abort?.addEventListener('abort', cleanup) - process.on('exit', cleanup) - - // Handle process exit - let resolveExit: () => void - let rejectExit: (error: Error) => void - const processExitPromise = new Promise((resolve, reject) => { - resolveExit = resolve - rejectExit = reject - }) + const handleAbort = () => { + void cleanup() + } + const handleProcessExit = () => { + void cleanup() + } + config.options?.abort?.addEventListener('abort', handleAbort) + process.on('exit', handleProcessExit) // Create query instance BEFORE registering close handler // to avoid temporal dependency on `query` variable const query = new Query(childStdin, child.stdout, processExitPromise, canCallTool) + if (typeof prompt !== 'string') { + void streamToStdin(prompt, child.stdin, config.options?.abort).catch(async (error) => { + const err = error instanceof Error ? error : new Error(String(error)) + if (!query.registerPromptFailure(err)) { + return + } + await Promise.race([ + cleanup(), + new Promise((resolve) => setTimeout(resolve, promptFailureCleanupTimeoutMs)) + ]) + query.setError(err) + rejectExit(err) + }) + } + // Register close handler - query is safely defined now child.on('close', (code) => { - if (config.options?.abort?.aborted) { + const promptFailure = query.getPromptFailure() + if (promptFailure) { + rejectExit(promptFailure) + } else if (config.options?.abort?.aborted) { const err = new AbortError('Claude Code process aborted by user') query.setError(err) rejectExit(err) @@ -418,7 +480,10 @@ export function query(config: { // Handle process errors child.on('error', (error) => { cleanupMcpConfig?.() - if (config.options?.abort?.aborted) { + const promptFailure = query.getPromptFailure() + if (promptFailure) { + rejectExit(promptFailure) + } else if (config.options?.abort?.aborted) { const err = new AbortError('Claude Code process aborted by user') query.setError(err) rejectExit(err) @@ -431,9 +496,9 @@ export function query(config: { // Cleanup on exit (catch rejection to avoid unhandled promise warning) processExitPromise.catch(() => {}).finally(() => { - cleanup() - process.removeListener('exit', cleanup) - config.options?.abort?.removeEventListener('abort', cleanup) + void cleanup() + process.removeListener('exit', handleProcessExit) + config.options?.abort?.removeEventListener('abort', handleAbort) if (process.env.CLAUDE_SDK_MCP_SERVERS) { delete process.env.CLAUDE_SDK_MCP_SERVERS } diff --git a/cli/src/claude/sdk/stream.test.ts b/cli/src/claude/sdk/stream.test.ts new file mode 100644 index 00000000..ee67d398 --- /dev/null +++ b/cli/src/claude/sdk/stream.test.ts @@ -0,0 +1,33 @@ +import { describe, expect, it } from 'vitest' +import { Stream } from './stream' + +describe('Stream', () => { + it('keeps the first error sticky even if done is called later', async () => { + const stream = new Stream() + const error = new Error('prompt failed') + + stream.error(error) + stream.done() + + await expect(stream.next()).rejects.toThrow('prompt failed') + }) + + it('ignores enqueue after terminal error', async () => { + const stream = new Stream() + const error = new Error('prompt failed') + + stream.error(error) + stream.enqueue('late-message') + + await expect(stream.next()).rejects.toThrow('prompt failed') + }) + + it('rejects a pending consumer when error arrives asynchronously', async () => { + const stream = new Stream() + const pending = stream.next() + + stream.error(new Error('prompt failed')) + + await expect(pending).rejects.toThrow('prompt failed') + }) +}) diff --git a/cli/src/claude/sdk/stream.ts b/cli/src/claude/sdk/stream.ts index 3d7d05eb..696929a9 100644 --- a/cli/src/claude/sdk/stream.ts +++ b/cli/src/claude/sdk/stream.ts @@ -12,7 +12,7 @@ export class Stream implements AsyncIterableIterator { private readResolve?: (value: IteratorResult) => void private readReject?: (error: Error) => void private isDone = false - private hasError?: Error + private terminalError?: Error private started = false constructor(private returned?: () => void) {} @@ -32,6 +32,10 @@ export class Stream implements AsyncIterableIterator { * Gets the next value from the stream */ async next(): Promise> { + if (this.terminalError) { + return Promise.reject(this.terminalError) + } + // Return queued items first if (this.queue.length > 0) { return Promise.resolve({ @@ -45,10 +49,6 @@ export class Stream implements AsyncIterableIterator { return Promise.resolve({ done: true, value: undefined }) } - if (this.hasError) { - return Promise.reject(this.hasError) - } - // Wait for new data return new Promise((resolve, reject) => { this.readResolve = resolve @@ -60,6 +60,10 @@ export class Stream implements AsyncIterableIterator { * Adds a value to the stream */ enqueue(value: T): void { + if (this.isDone || this.terminalError) { + return + } + if (this.readResolve) { // Direct delivery to waiting consumer const resolve = this.readResolve @@ -76,6 +80,10 @@ export class Stream implements AsyncIterableIterator { * Marks the stream as complete */ done(): void { + if (this.isDone || this.terminalError) { + return + } + this.isDone = true if (this.readResolve) { const resolve = this.readResolve @@ -89,7 +97,12 @@ export class Stream implements AsyncIterableIterator { * Propagates an error through the stream */ error(error: Error): void { - this.hasError = error + if (this.isDone || this.terminalError) { + return + } + + this.terminalError = error + this.queue = [] if (this.readReject) { const reject = this.readReject this.readResolve = undefined @@ -108,4 +121,8 @@ export class Stream implements AsyncIterableIterator { } return Promise.resolve({ done: true, value: undefined }) } -} \ No newline at end of file + + get hasTerminalError(): boolean { + return this.terminalError !== undefined + } +} diff --git a/cli/src/claude/sdk/types.ts b/cli/src/claude/sdk/types.ts index 9298237d..d48fab3d 100644 --- a/cli/src/claude/sdk/types.ts +++ b/cli/src/claude/sdk/types.ts @@ -175,6 +175,7 @@ export interface QueryOptions { settingsPath?: string strictMcpConfig?: boolean canCallTool?: CanCallToolCallback + promptFailureCleanupTimeoutMs?: number } /**