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",