From f91f5c837f32003dea8608add6eaff81d722e4ec Mon Sep 17 00:00:00 2001 From: Brian Love Date: Wed, 23 Sep 2026 10:24:03 -0700 Subject: [PATCH] fix(langgraph): honor cancellation for state writes Forward the existing AbortSignal to LangGraph state writes and cover cancellation through normal and protected transports with real HTTP regressions. Include the generated API documentation and parity inventory update. --- .../content/docs/langgraph/api/api-docs.json | 4 +- .../lib/transport/fetch-stream.transport.ts | 6 +- .../runtime/state-write-cancellation.spec.ts | 112 ++++++++++++++++++ scripts/react-parity/baseline.json | 10 +- 4 files changed, 123 insertions(+), 9 deletions(-) create mode 100644 libs/langgraph/src/runtime/state-write-cancellation.spec.ts diff --git a/apps/website/content/docs/langgraph/api/api-docs.json b/apps/website/content/docs/langgraph/api/api-docs.json index 17091a659..cf5b0820a 100644 --- a/apps/website/content/docs/langgraph/api/api-docs.json +++ b/apps/website/content/docs/langgraph/api/api-docs.json @@ -363,7 +363,7 @@ }, { "name": "updateState", - "signature": "updateState(threadId: string, values: Record, _signal: AbortSignal, options: object): Promise", + "signature": "updateState(threadId: string, values: Record, signal: AbortSignal, options: object): Promise", "description": "Update server-side thread state, e.g. to remove messages for regenerate rollback.", "params": [ { @@ -379,7 +379,7 @@ "optional": false }, { - "name": "_signal", + "name": "signal", "type": "AbortSignal", "description": "", "optional": false diff --git a/libs/langgraph/src/lib/transport/fetch-stream.transport.ts b/libs/langgraph/src/lib/transport/fetch-stream.transport.ts index 847a86da4..b88b976cf 100644 --- a/libs/langgraph/src/lib/transport/fetch-stream.transport.ts +++ b/libs/langgraph/src/lib/transport/fetch-stream.transport.ts @@ -193,17 +193,17 @@ export class FetchStreamTransport implements AgentTransport { async updateState( threadId: string, values: Record, - _signal: AbortSignal, + signal: AbortSignal, options?: { asNode?: string }, ): Promise { - const body: { values: Record; asNode?: string } = { values }; + const body: { values: Record; signal: AbortSignal; asNode?: string } = { values, signal }; if (options?.asNode !== undefined) { body.asNode = options.asNode; } try { await this.client.threads.updateState(threadId, body); } catch (error) { - this.rethrowOperationError(error, _signal); + this.rethrowOperationError(error, signal); } } diff --git a/libs/langgraph/src/runtime/state-write-cancellation.spec.ts b/libs/langgraph/src/runtime/state-write-cancellation.spec.ts new file mode 100644 index 000000000..034cfd172 --- /dev/null +++ b/libs/langgraph/src/runtime/state-write-cancellation.spec.ts @@ -0,0 +1,112 @@ +import { createServer, type ServerResponse } from 'node:http'; +import { afterEach, describe, expect, it, vi } from 'vitest'; +import { FetchStreamTransport } from '../lib/transport/fetch-stream.transport'; +import { deferred } from './testing/deferred'; + +const cleanup: (() => Promise)[] = []; +afterEach(async () => { + await Promise.all(cleanup.splice(0).map((close) => close())); +}); + +async function endpoint(hold = false) { + const started = deferred(); + const requests: { method?: string; url?: string; body: unknown }[] = []; + const server = createServer(async (request, response) => { + let body = ''; + for await (const chunk of request) body += chunk; + requests.push({ + method: request.method, + url: request.url, + body: JSON.parse(body), + }); + started.resolve(response); + if (!hold) { + response.setHeader('content-type', 'application/json'); + response.end(JSON.stringify({ checkpoint_id: 'saved' })); + } + }); + await new Promise((resolve) => server.listen(0, '127.0.0.1', resolve)); + cleanup.push( + () => + new Promise((resolve, reject) => { + server.closeAllConnections(); + server.close((error) => (error ? reject(error) : resolve())); + }) + ); + const address = server.address(); + if (!address || typeof address === 'string') + throw new Error('Expected TCP server'); + return { + url: `http://127.0.0.1:${address.port}`, + started: started.promise, + requests, + }; +} + +describe.each([false, true])( + 'state-write HTTP cancellation (protected=%s)', + (protectedErrors) => { + function transport(url: string) { + const report = vi.fn(); + return { + client: new FetchStreamTransport( + url, + undefined, + { maxRetries: 0 }, + protectedErrors ? report : undefined + ), + report, + }; + } + + it('does not send a state write when already aborted', async () => { + const server = await endpoint(); + const { client, report } = transport(server.url); + const controller = new AbortController(); + controller.abort(); + await expect( + client.updateState('thread-a', {}, controller.signal) + ).rejects.toMatchObject({ name: 'AbortError' }); + expect(server.requests).toEqual([]); + expect(report).not.toHaveBeenCalled(); + }); + + it('aborts an in-flight state write and closes its HTTP response', async () => { + const server = await endpoint(true); + const { client, report } = transport(server.url); + const controller = new AbortController(); + const write = client.updateState('thread-a', {}, controller.signal); + const rejected = expect(write).rejects.toMatchObject({ + name: 'AbortError', + }); + const response = await server.started; + const closed = new Promise((resolve) => + response.once('close', resolve) + ); + controller.abort(); + await rejected; + await closed; + expect(server.requests).toHaveLength(1); + expect(report).not.toHaveBeenCalled(); + }); + + it('preserves the state body and asNode for successful writes', async () => { + const server = await endpoint(); + const { client } = transport(server.url); + const values = Object.freeze({ messages: Object.freeze([]) }); + await client.updateState( + 'thread-a', + values, + new AbortController().signal, + { asNode: '__start__' } + ); + expect(server.requests).toEqual([ + { + method: 'POST', + url: '/threads/thread-a/state', + body: { values, as_node: '__start__' }, + }, + ]); + }); + } +); diff --git a/scripts/react-parity/baseline.json b/scripts/react-parity/baseline.json index bebab6c3b..f79aac36d 100644 --- a/scripts/react-parity/baseline.json +++ b/scripts/react-parity/baseline.json @@ -1,8 +1,10 @@ { "schemaVersion": 1, - "baselineHead": "82c5469743a8fc3484d395f03372dc6c6975e9b2", + "baselineHead": "ea15ba3e48b156872bed66066462fc0a30f0425a", "sourceState": { - "modified": [], + "modified": [ + "libs/langgraph/src/lib/transport/fetch-stream.transport.ts" + ], "untracked": [] }, "scope": { @@ -8356,7 +8358,7 @@ "path": "libs/langgraph/src/lib/transport/fetch-stream.transport.ts", "symbol": "FetchStreamTransport", "syntaxKind": "ClassDeclaration", - "signature": "export class FetchStreamTransport implements AgentTransport {\n private client: Client;\n private onThreadId?: (id: string) => void;\n private readonly protectErrors: boolean;\n private readonly reportOperationFailure?: RuntimeOperationFailureReporter;\n readonly protectsOperationErrors: boolean;\n constructor(apiUrl: string, onThreadId?: (id: string) => void, clientOptions?: LangGraphClientOptions, reportOperationFailure?: RuntimeOperationFailureReporter) {\n this.protectErrors = clientOptions?.defaultHeaders !== undefined || reportOperationFailure !== undefined;\n this.protectsOperationErrors = this.protectErrors;\n this.reportOperationFailure = reportOperationFailure;\n this.client = this.protectErrors\n ? ɵcreateProtectedLangGraphClient(apiUrl, clientOptions, createLangGraphRuntimeFetch(reportOperationFailure))\n : createLangGraphClient(apiUrl, clientOptions);\n this.onThreadId = onThreadId;\n }\n async *stream(assistantId: string, threadId: string | null, payload: unknown, signal: AbortSignal, options?: LangGraphSubmitOptions): AsyncIterable {\n let thread = threadId;\n if (!thread) {\n try {\n const t = await this.client.threads.create();\n thread = t.thread_id;\n }\n catch (error) {\n this.rethrowOperationError(error, signal);\n }\n try {\n this.onThreadId?.(thread);\n }\n catch (error) {\n this.rethrowLocalError(error, signal);\n }\n }\n let runPayload: ReturnType;\n try {\n runPayload = buildRunPayload(payload, signal, options);\n }\n catch (error) {\n return this.rethrowLocalError(error, signal);\n }\n let run: ReturnType;\n try {\n run = this.client.runs.stream(thread, assistantId, runPayload);\n }\n catch (error) {\n this.rethrowOperationError(error, signal);\n }\n yield* this.iterateSdkRun(run, signal);\n }\n async *joinStream(threadId: string, runId: string, lastEventId: string | undefined, signal: AbortSignal): AsyncIterable {\n let run: ReturnType;\n try {\n run = this.client.runs.joinStream(threadId, runId, {\n signal,\n ...(lastEventId !== undefined ? { lastEventId } : {}),\n });\n }\n catch (error) {\n this.rethrowOperationError(error, signal);\n }\n yield* this.iterateSdkRun(run, signal);\n }\n async getRunStatus(threadId: string, runId: string, signal: AbortSignal): Promise {\n try {\n const run = await this.client.runs.get(threadId, runId, { signal });\n if (run.run_id !== runId || run.thread_id !== threadId ||\n !['pending', 'running', 'success', 'error', 'timeout', 'interrupted'].includes(run.status))\n throw new Error('Invalid LangGraph run status response.');\n return run.status;\n }\n catch (error) {\n return this.rethrowOperationError(error, signal);\n }\n }\n async createQueuedRun(assistantId: string, threadId: string, payload: unknown, signal: AbortSignal, options?: LangGraphSubmitOptions): Promise {\n let runPayload: ReturnType & {\n multitaskStrategy: 'enqueue';\n };\n try {\n runPayload = {\n ...buildRunPayload(payload, signal, options),\n multitaskStrategy: 'enqueue',\n };\n }\n catch (error) {\n return this.rethrowLocalError(error, signal);\n }\n let run: Awaited>;\n try {\n run = await this.client.runs.create(threadId, assistantId, runPayload);\n }\n catch (error) {\n return this.rethrowOperationError(error, signal);\n }\n try {\n return {\n id: run.run_id,\n threadId: run.thread_id ?? threadId,\n values: payload,\n options: { multitaskStrategy: 'enqueue', signal },\n createdAt: run.created_at ? new Date(run.created_at) : new Date(),\n };\n }\n catch (error) {\n return this.rethrowLocalError(error, signal);\n }\n }\n async cancelRun(threadId: string, runId: string, signal: AbortSignal): Promise {\n try {\n await this.client.runs.cancel(threadId, runId, false, 'interrupt', { signal });\n }\n catch (error) {\n this.rethrowOperationError(error, signal);\n }\n }\n async getHistory(threadId: string, signal: AbortSignal): Promise {\n try {\n return await this.client.threads.getHistory(threadId, { signal });\n }\n catch (error) {\n return this.rethrowOperationError(error, signal);\n }\n }\n async updateState(threadId: string, values: Record, _signal: AbortSignal, options?: {\n asNode?: string;\n }): Promise {\n const body: {\n values: Record;\n asNode?: string;\n } = { values };\n if (options?.asNode !== undefined) {\n body.asNode = options.asNode;\n }\n try {\n await this.client.threads.updateState(threadId, body);\n }\n catch (error) {\n this.rethrowOperationError(error, _signal);\n }\n }\n private rethrowOperationError(error: unknown, signal: AbortSignal): never {\n if (!this.protectErrors)\n throw error;\n return projectLangGraphOperationFailure(error, signal, this.reportOperationFailure);\n }\n private rethrowLocalError(error: unknown, signal: AbortSignal): never {\n if (!this.protectErrors)\n throw error;\n return projectLangGraphOperationFailure(error, signal, undefined);\n }\n private async *iterateSdkRun(run: ReturnType | ReturnType, signal: AbortSignal): AsyncIterable {\n let iterator: AsyncIterator<{\n event: string;\n data: unknown;\n id?: unknown;\n }>;\n try {\n iterator = run[Symbol.asyncIterator]();\n }\n catch (error) {\n return this.rethrowOperationError(error, signal);\n }\n let completed = false;\n let failed = false;\n try {\n while (true) {\n let next: IteratorResult<{\n event: string;\n data: unknown;\n id?: unknown;\n }>;\n try {\n next = await iterator.next();\n }\n catch (error) {\n return this.rethrowOperationError(error, signal);\n }\n if (next.done) {\n completed = true;\n return;\n }\n try {\n yield { ...normalizeSdkEvent(next.value.event as StreamEvent['type'], next.value.data), sseId: typeof next.value.id === 'string' ? next.value.id : undefined };\n }\n catch (error) {\n return this.rethrowLocalError(error, signal);\n }\n }\n }\n catch (error) {\n failed = true;\n throw error;\n }\n finally {\n if (!completed) {\n try {\n await iterator.return?.();\n }\n catch (error) {\n if (!failed)\n this.rethrowOperationError(error, signal);\n }\n }\n }\n }\n}" + "signature": "export class FetchStreamTransport implements AgentTransport {\n private client: Client;\n private onThreadId?: (id: string) => void;\n private readonly protectErrors: boolean;\n private readonly reportOperationFailure?: RuntimeOperationFailureReporter;\n readonly protectsOperationErrors: boolean;\n constructor(apiUrl: string, onThreadId?: (id: string) => void, clientOptions?: LangGraphClientOptions, reportOperationFailure?: RuntimeOperationFailureReporter) {\n this.protectErrors = clientOptions?.defaultHeaders !== undefined || reportOperationFailure !== undefined;\n this.protectsOperationErrors = this.protectErrors;\n this.reportOperationFailure = reportOperationFailure;\n this.client = this.protectErrors\n ? ɵcreateProtectedLangGraphClient(apiUrl, clientOptions, createLangGraphRuntimeFetch(reportOperationFailure))\n : createLangGraphClient(apiUrl, clientOptions);\n this.onThreadId = onThreadId;\n }\n async *stream(assistantId: string, threadId: string | null, payload: unknown, signal: AbortSignal, options?: LangGraphSubmitOptions): AsyncIterable {\n let thread = threadId;\n if (!thread) {\n try {\n const t = await this.client.threads.create();\n thread = t.thread_id;\n }\n catch (error) {\n this.rethrowOperationError(error, signal);\n }\n try {\n this.onThreadId?.(thread);\n }\n catch (error) {\n this.rethrowLocalError(error, signal);\n }\n }\n let runPayload: ReturnType;\n try {\n runPayload = buildRunPayload(payload, signal, options);\n }\n catch (error) {\n return this.rethrowLocalError(error, signal);\n }\n let run: ReturnType;\n try {\n run = this.client.runs.stream(thread, assistantId, runPayload);\n }\n catch (error) {\n this.rethrowOperationError(error, signal);\n }\n yield* this.iterateSdkRun(run, signal);\n }\n async *joinStream(threadId: string, runId: string, lastEventId: string | undefined, signal: AbortSignal): AsyncIterable {\n let run: ReturnType;\n try {\n run = this.client.runs.joinStream(threadId, runId, {\n signal,\n ...(lastEventId !== undefined ? { lastEventId } : {}),\n });\n }\n catch (error) {\n this.rethrowOperationError(error, signal);\n }\n yield* this.iterateSdkRun(run, signal);\n }\n async getRunStatus(threadId: string, runId: string, signal: AbortSignal): Promise {\n try {\n const run = await this.client.runs.get(threadId, runId, { signal });\n if (run.run_id !== runId || run.thread_id !== threadId ||\n !['pending', 'running', 'success', 'error', 'timeout', 'interrupted'].includes(run.status))\n throw new Error('Invalid LangGraph run status response.');\n return run.status;\n }\n catch (error) {\n return this.rethrowOperationError(error, signal);\n }\n }\n async createQueuedRun(assistantId: string, threadId: string, payload: unknown, signal: AbortSignal, options?: LangGraphSubmitOptions): Promise {\n let runPayload: ReturnType & {\n multitaskStrategy: 'enqueue';\n };\n try {\n runPayload = {\n ...buildRunPayload(payload, signal, options),\n multitaskStrategy: 'enqueue',\n };\n }\n catch (error) {\n return this.rethrowLocalError(error, signal);\n }\n let run: Awaited>;\n try {\n run = await this.client.runs.create(threadId, assistantId, runPayload);\n }\n catch (error) {\n return this.rethrowOperationError(error, signal);\n }\n try {\n return {\n id: run.run_id,\n threadId: run.thread_id ?? threadId,\n values: payload,\n options: { multitaskStrategy: 'enqueue', signal },\n createdAt: run.created_at ? new Date(run.created_at) : new Date(),\n };\n }\n catch (error) {\n return this.rethrowLocalError(error, signal);\n }\n }\n async cancelRun(threadId: string, runId: string, signal: AbortSignal): Promise {\n try {\n await this.client.runs.cancel(threadId, runId, false, 'interrupt', { signal });\n }\n catch (error) {\n this.rethrowOperationError(error, signal);\n }\n }\n async getHistory(threadId: string, signal: AbortSignal): Promise {\n try {\n return await this.client.threads.getHistory(threadId, { signal });\n }\n catch (error) {\n return this.rethrowOperationError(error, signal);\n }\n }\n async updateState(threadId: string, values: Record, signal: AbortSignal, options?: {\n asNode?: string;\n }): Promise {\n const body: {\n values: Record;\n signal: AbortSignal;\n asNode?: string;\n } = { values, signal };\n if (options?.asNode !== undefined) {\n body.asNode = options.asNode;\n }\n try {\n await this.client.threads.updateState(threadId, body);\n }\n catch (error) {\n this.rethrowOperationError(error, signal);\n }\n }\n private rethrowOperationError(error: unknown, signal: AbortSignal): never {\n if (!this.protectErrors)\n throw error;\n return projectLangGraphOperationFailure(error, signal, this.reportOperationFailure);\n }\n private rethrowLocalError(error: unknown, signal: AbortSignal): never {\n if (!this.protectErrors)\n throw error;\n return projectLangGraphOperationFailure(error, signal, undefined);\n }\n private async *iterateSdkRun(run: ReturnType | ReturnType, signal: AbortSignal): AsyncIterable {\n let iterator: AsyncIterator<{\n event: string;\n data: unknown;\n id?: unknown;\n }>;\n try {\n iterator = run[Symbol.asyncIterator]();\n }\n catch (error) {\n return this.rethrowOperationError(error, signal);\n }\n let completed = false;\n let failed = false;\n try {\n while (true) {\n let next: IteratorResult<{\n event: string;\n data: unknown;\n id?: unknown;\n }>;\n try {\n next = await iterator.next();\n }\n catch (error) {\n return this.rethrowOperationError(error, signal);\n }\n if (next.done) {\n completed = true;\n return;\n }\n try {\n yield { ...normalizeSdkEvent(next.value.event as StreamEvent['type'], next.value.data), sseId: typeof next.value.id === 'string' ? next.value.id : undefined };\n }\n catch (error) {\n return this.rethrowLocalError(error, signal);\n }\n }\n }\n catch (error) {\n failed = true;\n throw error;\n }\n finally {\n if (!completed) {\n try {\n await iterator.return?.();\n }\n catch (error) {\n if (!failed)\n this.rethrowOperationError(error, signal);\n }\n }\n }\n }\n}" } ] }, @@ -12534,7 +12536,7 @@ "id": "source:libs/langgraph/src/lib/transport/fetch-stream.transport.ts", "kind": "source", "path": "libs/langgraph/src/lib/transport/fetch-stream.transport.ts", - "sha256": "a4a66891600fdf1bed6cc769deca86b87a724f392568369ba006c987fc583f95" + "sha256": "2b899f3fb3758f8c57c500b902d16d826fbfd1090c3b160692f0b01456024981" }, { "id": "source:libs/langgraph/src/lib/transport/mock-stream.transport.ts",