@dbos-inc/vercel-ai 0.4.4 → 0.5.4

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -1,4 +1,4 @@
1
- import { randomUUID } from 'node:crypto';
1
+ import { createHash, randomUUID } from 'node:crypto';
2
2
  import { DBOS, Error as DBOSErrors, StatusString } from '@dbos-inc/dbos-sdk';
3
3
  import type { UIMessageChunk } from 'ai' with { 'resolution-mode': 'import' };
4
4
  import type { LanguageModelV4FinishReason, LanguageModelV4StreamPart } from '@ai-sdk/provider' with { 'resolution-mode': 'import' };
@@ -10,7 +10,7 @@ export type DurableStreamOptions = string | { key: string; maxBatchParts?: numbe
10
10
  export type DurableStreamRecord =
11
11
  | { kind: 'model'; step: number; attempt: string; parts: LanguageModelV4StreamPart[] }
12
12
  | { kind: 'model-end'; step: number; attempt: string; finishReason?: LanguageModelV4FinishReason; aborted?: true }
13
- | { kind: 'tool'; step: number; attempt: number; toolCallId: string; output?: unknown; errorText?: string }
13
+ | { kind: 'tool'; step: number; attempt: number; toolCallId: string; output?: unknown; errorText?: string; chunks?: UIMessageChunk[] }
14
14
  | { kind: 'ui'; step?: number; attempt?: number; chunks: UIMessageChunk[] }
15
15
  | { kind: 'end'; finishReason: string };
16
16
 
@@ -123,9 +123,14 @@ export class ModelStreamWriter {
123
123
  }
124
124
  }
125
125
 
126
- /** Records a tool call's outcome from inside its step. */
127
- export function writeToolRecord(key: string, toolCallId: string, outcome: { output: unknown } | { errorText: string }): Promise<void> {
128
- const record: DurableStreamRecord = { kind: 'tool', ...stepInfo(), toolCallId, ...outcome };
126
+ /** Records a tool call's outcome, with the message chunks it wrote, from inside its step. */
127
+ export function writeToolRecord(
128
+ key: string,
129
+ toolCallId: string,
130
+ outcome: { output: unknown } | { errorText: string },
131
+ chunks: UIMessageChunk[] = [],
132
+ ): Promise<void> {
133
+ const record: DurableStreamRecord = { kind: 'tool', ...stepInfo(), toolCallId, ...outcome, ...(chunks.length > 0 && { chunks }) };
129
134
  return DBOS.writeStream(key, record);
130
135
  }
131
136
 
@@ -158,7 +163,10 @@ export interface ReadDurableStreamOptions {
158
163
  key: string;
159
164
  /** Id for the `start` chunk; omitted on a resume (`offset` > 0). */
160
165
  messageId?: string;
161
- /** Number of records already consumed, from the last `data-dbos-offset` chunk. */
166
+ /**
167
+ * Number of records already consumed, from the last `data-dbos-offset` chunk. A resume from mid-text can't be applied
168
+ * to a partial message by the AI SDK (it has no record of the open part); to rebuild a message, read from 0 instead.
169
+ */
162
170
  offset?: number;
163
171
  /** Defaults to `DBOS`; pass a `DBOSClient` to read from a process that has not launched DBOS. */
164
172
  client?: DurableStreamSource;
@@ -190,11 +198,13 @@ async function* uiChunks(options: ReadDurableStreamOptions): AsyncGenerator<UIMe
190
198
  const client: DurableStreamSource = options.client ?? DBOS;
191
199
  const state: ReaderState = {
192
200
  offset: options.offset ?? 0,
193
- resumed: (options.offset ?? 0) > 0,
194
201
  openParts: new Map(),
202
+ detachedParts: new Set(),
203
+ toolChunksSent: new Map(),
195
204
  ended: false,
196
205
  };
197
206
  if (state.offset === 0) yield { type: 'start', messageId: options.messageId };
207
+ else await restoreOpenCall(client, workflowID, key, state);
198
208
  const emit = (record: DurableStreamRecord) => emitRecord(state, record, { sendReasoning, sendSources, onError });
199
209
 
200
210
  // Phase 1: everything already stored, one value per query until an offset is empty; a superseded attempt is skipped whole.
@@ -208,12 +218,16 @@ async function* uiChunks(options: ReadDurableStreamOptions): AsyncGenerator<UIMe
208
218
  }
209
219
  }
210
220
  const finalAttempt = new Map<number, string>();
211
- for (const record of history) {
221
+ // A re-executed tool call writes another record; the last one matches its checkpoint, so only its chunks (if any) are shown.
222
+ const finalToolRecord = new Map<string, number>();
223
+ history.forEach((record, index) => {
212
224
  if (record.kind === 'model' || record.kind === 'model-end') finalAttempt.set(record.step, record.attempt);
213
- }
214
- for (const record of history) {
225
+ if (record.kind === 'tool') finalToolRecord.set(toolKey(record), index);
226
+ });
227
+ for (const [index, record] of history.entries()) {
215
228
  const stale = (record.kind === 'model' || record.kind === 'model-end') && finalAttempt.get(record.step) !== record.attempt;
216
- yield* stale ? skipRecord(state) : emit(record);
229
+ const staleChunks = record.kind === 'tool' && record.chunks !== undefined && finalToolRecord.get(toolKey(record)) !== index;
230
+ yield* stale ? skipRecord(state) : emit(staleChunks ? { ...record, chunks: undefined } : record);
217
231
  if (state.ended) return;
218
232
  }
219
233
 
@@ -240,15 +254,23 @@ async function* uiChunks(options: ReadDurableStreamOptions): AsyncGenerator<UIMe
240
254
 
241
255
  interface ReaderState {
242
256
  offset: number;
243
- resumed: boolean;
244
257
  openStep?: number;
245
258
  openAttempt?: string;
246
259
  // Text/reasoning parts of the open attempt that have started but not ended, by UI part id.
247
260
  openParts: Map<string, 'text' | 'reasoning'>;
261
+ // Open parts from before a resume: the client's resumed stream no longer tracks them, so they are never ended.
262
+ detachedParts: Set<string>;
263
+ // Digest of the chunks emitted for each tool call (by toolKey).
264
+ toolChunksSent: Map<string, string>;
248
265
  finishReason?: string;
249
266
  ended: boolean;
250
267
  }
251
268
 
269
+ // A step's function id survives re-execution, and with it tell apart calls that reuse a provider's tool call id.
270
+ function toolKey(record: Extract<DurableStreamRecord, { kind: 'tool' }>): string {
271
+ return `${record.step}:${record.toolCallId}`;
272
+ }
273
+
252
274
  function offsetChunk(state: ReaderState): UIMessageChunk {
253
275
  return { type: 'data-dbos-offset', data: { offset: state.offset }, transient: true } as UIMessageChunk;
254
276
  }
@@ -263,17 +285,90 @@ function* closeStep(state: ReaderState): Generator<UIMessageChunk> {
263
285
  state.openStep = undefined;
264
286
  state.openAttempt = undefined;
265
287
  state.openParts.clear();
288
+ state.detachedParts.clear();
289
+ }
290
+
291
+ // Records outside the open model call a resume reads past before giving up; a reconnect mid-call finds it at once.
292
+ const RESTORE_SCAN_LIMIT = 100;
293
+
294
+ /**
295
+ * On resume, rebuild the model call that was open at `offset` by reading back to its start, so the reader continues
296
+ * exactly as an uninterrupted one would: a re-executed attempt is superseded, and step boundaries fall in the same place.
297
+ * Past RESTORE_SCAN_LIMIT unrelated records it gives up, and the resumed stream only lacks that step's `finish-step`.
298
+ */
299
+ async function restoreOpenCall(client: DurableStreamSource, workflowID: string, key: string, state: ReaderState): Promise<void> {
300
+ let step: number | undefined;
301
+ let attempt: string | undefined;
302
+ // Set once an earlier step's record is reached: the open call is fully read, and only a finish reason is still wanted.
303
+ let callRead = false;
304
+ let budget = RESTORE_SCAN_LIMIT;
305
+ const ended = new Set<string>();
306
+ for (let index = state.offset - 1; index >= 0; index--) {
307
+ let record: DurableStreamRecord;
308
+ try {
309
+ record = await client.readStreamOffset<DurableStreamRecord>(workflowID, key, index, { timeoutSeconds: 0 });
310
+ } catch (error) {
311
+ // A real offset never passes the stored records, so a gap means there is nothing to restore.
312
+ if (DBOSErrors.isStreamTimeoutError(error)) break;
313
+ throw error;
314
+ }
315
+ if (record.kind === 'end') break;
316
+ const model = record.kind === 'model' || record.kind === 'model-end' ? record : undefined;
317
+ if (model && step === undefined) {
318
+ step = model.step;
319
+ attempt = model.attempt;
320
+ }
321
+ if (model && !callRead && model.step === step) {
322
+ if (model.attempt === attempt) restorePart(state, ended, model);
323
+ continue;
324
+ }
325
+ if (model) callRead = true;
326
+ // As in an uninterrupted read, the finish reason carries over from the latest earlier call that ended.
327
+ if (callRead && (state.finishReason !== undefined || model?.kind === 'model-end')) {
328
+ if (model?.kind === 'model-end') state.finishReason ??= finishReasonOf(model);
329
+ break;
330
+ }
331
+ if (--budget === 0) break;
332
+ }
333
+ if (step === undefined) return;
334
+ state.openStep = step;
335
+ state.openAttempt = attempt;
336
+ }
337
+
338
+ function finishReasonOf(record: Extract<DurableStreamRecord, { kind: 'model-end' }>): string | undefined {
339
+ return record.aborted ? 'other' : record.finishReason?.unified;
340
+ }
341
+
342
+ // One record of the open call, read backward: a part is open if its start comes before (i.e. is read after) any end.
343
+ function restorePart(state: ReaderState, ended: Set<string>, record: Extract<DurableStreamRecord, { kind: 'model' | 'model-end' }>): void {
344
+ if (record.kind === 'model-end') {
345
+ state.finishReason ??= finishReasonOf(record);
346
+ return;
347
+ }
348
+ for (const part of [...record.parts].reverse()) {
349
+ const kind = part.type === 'text-start' || part.type === 'text-end' ? 'text' : part.type === 'reasoning-start' || part.type === 'reasoning-end' ? 'reasoning' : undefined;
350
+ if (kind === undefined || !('id' in part)) continue;
351
+ const id = `${record.attempt}:${part.id}`;
352
+ if (part.type.endsWith('-end')) ended.add(id);
353
+ else if (!ended.has(id)) {
354
+ state.openParts.set(id, kind);
355
+ state.detachedParts.add(id);
356
+ }
357
+ }
266
358
  }
267
359
 
268
360
  // A live re-execution of the open step: end the stale attempt's parts and tell the client which ones to discard.
269
361
  function* supersede(state: ReaderState, attempt: string): Generator<UIMessageChunk> {
270
- for (const [id, kind] of state.openParts) yield { type: kind === 'text' ? 'text-end' : 'reasoning-end', id };
362
+ for (const [id, kind] of state.openParts) {
363
+ if (!state.detachedParts.has(id)) yield { type: kind === 'text' ? 'text-end' : 'reasoning-end', id };
364
+ }
271
365
  yield {
272
366
  type: 'data-dbos-superseded',
273
367
  data: { attempt: state.openAttempt, parts: [...state.openParts.keys()] },
274
368
  transient: true,
275
369
  } as UIMessageChunk;
276
370
  state.openParts.clear();
371
+ state.detachedParts.clear();
277
372
  state.openAttempt = attempt;
278
373
  }
279
374
 
@@ -287,8 +382,7 @@ function* emitRecord(
287
382
  case 'model': {
288
383
  if (state.openStep !== record.step) {
289
384
  yield* closeStep(state);
290
- if (!state.resumed) yield { type: 'start-step' };
291
- state.resumed = false;
385
+ yield { type: 'start-step' };
292
386
  state.openStep = record.step;
293
387
  state.openAttempt = record.attempt;
294
388
  } else if (state.openAttempt !== record.attempt) {
@@ -307,14 +401,27 @@ function* emitRecord(
307
401
  }
308
402
  case 'model-end':
309
403
  // The stream outlives the call: the workflow may run more calls, so only its end (or closeDurableStream) ends the turn.
310
- if (record.attempt === state.openAttempt) state.finishReason = record.aborted ? 'other' : record.finishReason?.unified;
404
+ if (record.attempt === state.openAttempt) state.finishReason = finishReasonOf(record);
311
405
  break;
312
- case 'tool':
406
+ case 'tool': {
407
+ const key = toolKey(record);
408
+ const chunks = record.chunks ?? [];
409
+ const sent = state.toolChunksSent.get(key);
410
+ // A live re-execution that wrote the same chunks changes nothing; different ones (or none) replace what the client has.
411
+ if (chunks.length > 0 || sent !== undefined) {
412
+ const digest = createHash('sha256').update(JSON.stringify(chunks)).digest('base64');
413
+ if (sent !== digest) {
414
+ if (sent !== undefined) yield { type: 'data-dbos-tool-superseded', data: { toolCallId: record.toolCallId }, transient: true } as UIMessageChunk;
415
+ state.toolChunksSent.set(key, digest);
416
+ yield* chunks;
417
+ }
418
+ }
313
419
  // Local tool errors are masked like the AI SDK does; provider-executed ones (in model records) pass through verbatim.
314
420
  yield record.errorText !== undefined
315
421
  ? { type: 'tool-output-error', toolCallId: record.toolCallId, errorText: filter.onError(new Error(record.errorText)) }
316
422
  : { type: 'tool-output-available', toolCallId: record.toolCallId, output: record.output };
317
423
  break;
424
+ }
318
425
  case 'ui':
319
426
  yield* record.chunks;
320
427
  break;
package/src/index.ts CHANGED
@@ -1,6 +1,6 @@
1
1
  export { durableCalls, DurableCallsOptions, durableEmbeddingCalls, durableImageCalls } from './middleware';
2
2
  export { durableMCPTools, DurableMCPToolsOptions, MCPClientLike } from './mcp';
3
- export { durableTools, DurableToolsOptions } from './tools';
3
+ export { durableTools, DurableToolsOptions, toolWriter } from './tools';
4
4
  export {
5
5
  closeDurableStream,
6
6
  DurableStreamOptions,
package/src/internal.ts CHANGED
@@ -21,6 +21,31 @@ export function runDurableStep<T>(name: string, fn: () => Promise<T>, config: St
21
21
  });
22
22
  }
23
23
 
24
+ // DBOS 5.1+ fires stepStatus.cancelSignal when the step's workflow is cancelled; older versions have none.
25
+ export function stepCancelSignal(): AbortSignal | undefined {
26
+ return (DBOS.stepStatus as { cancelSignal?: AbortSignal } | undefined)?.cancelSignal;
27
+ }
28
+
29
+ /**
30
+ * DBOS discards a timed-out attempt's result, so a stream record of it would contradict the checkpoint. Returns the
31
+ * outcome to record instead: the timeout when the step ends with it, `null` when a retry follows, `undefined` if not timed out.
32
+ */
33
+ export async function timedOutOutcome(shouldRetry: StepConfig['shouldRetry']): Promise<{ errorText: string } | null | undefined> {
34
+ const status = DBOS.stepStatus;
35
+ if (status?.timeoutSignal?.aborted !== true) return undefined;
36
+ const reason: unknown = status.timeoutSignal.reason;
37
+ const lastAttempt = status.currentAttempt === undefined || status.currentAttempt >= (status.maxAttempts ?? 1);
38
+ // DBOS asks the step's shouldRetry about this same error (the signal's reason) before retrying; a throw ends the step.
39
+ const retried = !lastAttempt && (await Promise.resolve(shouldRetry ? shouldRetry(reason) : true).catch(() => false));
40
+ return retried ? null : { errorText: reason instanceof Error ? reason.message : 'The step timed out.' };
41
+ }
42
+
43
+ // Fires when any given signal does; a lone signal is returned as-is.
44
+ export function anySignal(...signals: (AbortSignal | undefined)[]): AbortSignal | undefined {
45
+ const defined = signals.filter((s): s is AbortSignal => s !== undefined);
46
+ return defined.length <= 1 ? defined[0] : AbortSignal.any(defined);
47
+ }
48
+
24
49
  export function isAsyncIterable(value: unknown): value is AsyncIterable<unknown> {
25
50
  return typeof (value as AsyncIterable<unknown> | null | undefined)?.[Symbol.asyncIterator] === 'function';
26
51
  }
package/src/mcp.ts CHANGED
@@ -1,6 +1,6 @@
1
- import { StepConfig } from '@dbos-inc/dbos-sdk';
1
+ import { DBOS, StepConfig } from '@dbos-inc/dbos-sdk';
2
2
  import type { ToolSet } from 'ai' with { 'resolution-mode': 'import' };
3
- import { isAsyncIterable, runDurableStep, withErrorClassification } from './internal';
3
+ import { anySignal, isAsyncIterable, runDurableStep, stepCancelSignal, timedOutOutcome, withErrorClassification } from './internal';
4
4
  import { writeToolRecord } from './durable-stream';
5
5
 
6
6
  // Structural type for an MCP client (e.g. from @ai-sdk/mcp) — deliberately loose: the AI SDK ecosystem
@@ -112,11 +112,20 @@ export async function durableMCPTools(client: MCPClientLike, options: DurableMCP
112
112
  return run(
113
113
  `mcp.tool.${name}.${toolCallId ?? 'call'}`,
114
114
  async () => {
115
+ // Stop the call when the attempt times out or the workflow is cancelled, as well as on the caller's abort.
116
+ const cancelSignal = stepCancelSignal();
117
+ const record = async (outcome: { output: unknown } | { errorText: string }) => {
118
+ if (!durableStream || !toolCallId) return;
119
+ const timeout = await timedOutOutcome(callConfig.shouldRetry);
120
+ if (timeout === null) return;
121
+ await writeToolRecord(durableStream, toolCallId, timeout ?? outcome);
122
+ };
123
+ const abortSignal = anySignal(signal, DBOS.stepStatus?.timeoutSignal, cancelSignal);
115
124
  let output: unknown;
116
125
  try {
117
126
  const tool = (await client.tools(toolOptions))[name] as MCPToolLike | undefined;
118
127
  if (typeof tool?.execute !== 'function') throw new Error(`MCP tool "${name}" is not executable.`);
119
- output = await tool.execute(input, execOptions);
128
+ output = await tool.execute(input, abortSignal === signal ? execOptions : { ...execOptions, abortSignal });
120
129
  // A streaming execute can't checkpoint mid-flight; drain it and record the final value (the last yield).
121
130
  if (isAsyncIterable(output)) {
122
131
  let last: unknown;
@@ -124,12 +133,10 @@ export async function durableMCPTools(client: MCPClientLike, options: DurableMCP
124
133
  output = last;
125
134
  }
126
135
  } catch (error) {
127
- if (durableStream && toolCallId) {
128
- await writeToolRecord(durableStream, toolCallId, { errorText: error instanceof Error ? error.message : String(error) });
129
- }
136
+ if (!cancelSignal?.aborted) await record({ errorText: error instanceof Error ? error.message : String(error) });
130
137
  throw error;
131
138
  }
132
- if (durableStream && toolCallId) await writeToolRecord(durableStream, toolCallId, { output });
139
+ await record({ output });
133
140
  return output;
134
141
  },
135
142
  callConfig,
package/src/middleware.ts CHANGED
@@ -20,7 +20,7 @@ import type {
20
20
  SharedV4ProviderMetadata,
21
21
  SharedV4Warning,
22
22
  } from '@ai-sdk/provider' with { 'resolution-mode': 'import' };
23
- import { assertNotInTransaction, isInWorkflowFunction, restoreAISDKErrorIdentity, withErrorClassification } from './internal';
23
+ import { anySignal, assertNotInTransaction, isInWorkflowFunction, restoreAISDKErrorIdentity, stepCancelSignal, withErrorClassification } from './internal';
24
24
  import { type DurableStreamOptions, ModelStreamWriter, resolveDurableStream } from './durable-stream';
25
25
 
26
26
  export interface DurableCallsOptions extends StepConfig {
@@ -75,7 +75,10 @@ export function durableCalls(options: DurableCallsOptions = {}): LanguageModelMi
75
75
  try {
76
76
  return await DBOS.runStep(
77
77
  async () => {
78
- const result = ensureResponseMetadata(encodeBinaryContent(omitBodies(await doGenerate(), include)));
78
+ // The provider call also stops when the attempt times out or the workflow is cancelled.
79
+ const signal = anySignal(params.abortSignal, DBOS.stepStatus?.timeoutSignal, stepCancelSignal());
80
+ const raw = signal === params.abortSignal ? await doGenerate() : await model.doGenerate({ ...params, abortSignal: signal });
81
+ const result = ensureResponseMetadata(encodeBinaryContent(omitBodies(raw, include)));
79
82
  // A non-streaming call writes its whole output at once, so the stream holds every call the loop makes, not only streamed ones.
80
83
  if (streamConfig) {
81
84
  const streamWriter = new ModelStreamWriter(streamConfig);
@@ -157,20 +160,26 @@ export function durableCalls(options: DurableCallsOptions = {}): LanguageModelMi
157
160
  const accumulator = new StreamAccumulator();
158
161
  // A timed-out attempt is abandoned by DBOS (its outcome is discarded) but keeps running; stop it so it can't emit alongside a retry.
159
162
  const timeoutSignal = DBOS.stepStatus?.timeoutSignal;
163
+ // A workflow cancel is not a consumer abort: it must fail the attempt so nothing is checkpointed and a resume re-runs the call.
164
+ const cancelSignal = stepCancelSignal();
165
+ const cancelledWorkflow = () => cancelSignal?.aborted === true;
166
+ const providerSignal = anySignal(abortSignal, timeoutSignal, cancelSignal);
160
167
  let reader: ReadableStreamDefaultReader<LanguageModelV4StreamPart> | undefined;
161
168
  let sawFinish = false;
162
169
  const abandon = () => void reader?.cancel().catch(() => {});
163
170
  timeoutSignal?.addEventListener('abort', abandon, { once: true });
171
+ cancelSignal?.addEventListener('abort', abandon, { once: true });
164
172
  // A consumer abort tears the provider call down too; the attempt is then recorded as aborted below.
165
173
  abortSignal?.addEventListener('abort', abandon, { once: true });
166
174
  // Step-scope stream writes: cheap, replay-safe, and never duplicated since a retry is refused once content has streamed.
167
175
  const streamWriter = streamConfig ? new ModelStreamWriter(streamConfig) : undefined;
168
176
  try {
169
- reader = (await doStream()).stream.getReader();
177
+ const result = providerSignal === abortSignal ? await doStream() : await model.doStream({ ...params, abortSignal: providerSignal });
178
+ reader = result.stream.getReader();
170
179
  for (;;) {
171
180
  const { done, value: part } = await reader.read();
172
181
  if (timeoutSignal?.aborted) throw (timeoutSignal.reason ?? new Error('step attempt timed out'));
173
- if (done || aborted()) break;
182
+ if (done || aborted() || cancelledWorkflow()) break;
174
183
  if (part.type === 'error') {
175
184
  // A cancelled consumer abandoned this call; don't let a late failure become the step outcome, or replay would fail where the live run succeeded.
176
185
  if (cancelled) break;
@@ -183,21 +192,26 @@ export function durableCalls(options: DurableCallsOptions = {}): LanguageModelMi
183
192
  accumulator.add(part);
184
193
  }
185
194
  // No terminal part and no output: fail (retryably) like the AI SDK's NoOutputGeneratedError, instead of checkpointing a permanent empty success.
186
- if (!sawFinish && !accumulator.hasContent && !cancelled && !aborted()) {
195
+ if (!sawFinish && !accumulator.hasContent && !cancelled && !aborted() && !cancelledWorkflow()) {
187
196
  throw new Error('Model stream ended without a finish part or any output.');
188
197
  }
189
198
  } catch (error) {
190
- // Same rule for stream-level failures (doStream or a read rejecting) after a cancel; after an abort, the abort is the outcome.
191
- if (!cancelled && !aborted()) {
199
+ // Same rule for stream-level failures (doStream or a read rejecting) after a cancel; after an abort or a workflow cancel, that is the outcome.
200
+ if (!cancelled && !aborted() && !cancelledWorkflow()) {
192
201
  await streamWriter?.abandon();
193
202
  throw error;
194
203
  }
195
204
  } finally {
196
205
  timeoutSignal?.removeEventListener('abort', abandon);
206
+ cancelSignal?.removeEventListener('abort', abandon);
197
207
  abortSignal?.removeEventListener('abort', abandon);
198
208
  // Tear down the provider stream on early exits (error part, post-cancel break); a no-op after a clean drain.
199
209
  void reader?.cancel().catch(() => {});
200
210
  }
211
+ if (cancelledWorkflow()) {
212
+ await streamWriter?.abandon();
213
+ throw cancelSignal!.reason;
214
+ }
201
215
  // Record the abort as the step's failure; replay rethrows it, so the workflow must catch aborts it means to survive.
202
216
  if (aborted()) {
203
217
  await streamWriter?.end({ aborted: true });
@@ -272,14 +286,17 @@ export function durableEmbeddingCalls(options: StepConfig = {}): EmbeddingModelM
272
286
  then: (onfulfilled, onrejected) =>
273
287
  Promise.resolve(isInWorkflowFunction() ? false : model.supportsParallelCalls).then(onfulfilled, onrejected),
274
288
  }),
275
- wrapEmbed: async ({ doEmbed, model }) => {
289
+ wrapEmbed: async ({ doEmbed, params, model }) => {
276
290
  assertNotInTransaction('embed');
277
291
  if (!isInWorkflowFunction()) {
278
292
  return await doEmbed();
279
293
  }
280
294
  const workflowID = await enterDurableModelCall();
281
295
  try {
282
- return await DBOS.runStep(async () => doEmbed(), {
296
+ return await DBOS.runStep(async () => {
297
+ const signal = anySignal(params.abortSignal, DBOS.stepStatus?.timeoutSignal, stepCancelSignal());
298
+ return signal === params.abortSignal ? doEmbed() : model.doEmbed({ ...params, abortSignal: signal });
299
+ }, {
283
300
  ...stepConfig,
284
301
  name: stepConfig.name ?? stepName(model, 'embed'),
285
302
  });
@@ -301,13 +318,16 @@ export function durableImageCalls(options: StepConfig = {}): ImageModelMiddlewar
301
318
  const stepConfig = withErrorClassification(options);
302
319
  return {
303
320
  specificationVersion: 'v4',
304
- wrapGenerate: async ({ doGenerate, model }) => {
321
+ wrapGenerate: async ({ doGenerate, params, model }) => {
305
322
  assertNotInTransaction('generateImage');
306
323
  if (!isInWorkflowFunction()) {
307
324
  return await doGenerate();
308
325
  }
309
326
  try {
310
- return await DBOS.runStep(async () => encodeImageResult(await doGenerate()), {
327
+ return await DBOS.runStep(async () => {
328
+ const signal = anySignal(params.abortSignal, DBOS.stepStatus?.timeoutSignal, stepCancelSignal());
329
+ return encodeImageResult(signal === params.abortSignal ? await doGenerate() : await model.doGenerate({ ...params, abortSignal: signal }));
330
+ }, {
311
331
  ...stepConfig,
312
332
  name: stepConfig.name ?? stepName(model, 'image'),
313
333
  });
package/src/tools.ts CHANGED
@@ -1,14 +1,108 @@
1
+ import { AsyncLocalStorage } from 'node:async_hooks';
1
2
  import { DBOS, StepConfig } from '@dbos-inc/dbos-sdk';
2
- import type { ToolSet } from 'ai' with { 'resolution-mode': 'import' };
3
- import { assertNotInTransaction, isAsyncIterable, isInWorkflowFunction, runDurableStep, withErrorClassification } from './internal';
3
+ import type { ToolSet, UIMessageChunk, UIMessageStreamWriter } from 'ai' with { 'resolution-mode': 'import' };
4
+ import { anySignal, assertNotInTransaction, isAsyncIterable, isInWorkflowFunction, runDurableStep, stepCancelSignal, timedOutOutcome, withErrorClassification } from './internal';
4
5
  import { AGENT_TOOL } from './agent-tool';
5
- import { writeToolRecord } from './durable-stream';
6
+ import { writeDurableStream, writeToolRecord } from './durable-stream';
6
7
 
7
8
  export interface DurableToolsOptions extends StepConfig {
8
9
  /** Per-tool step config overriding the defaults; `false` leaves that tool non-durable. */
9
10
  tools?: Record<string, StepConfig | false>;
10
11
  /** Write each tool call's output (or error) to this durable stream from inside its step. */
11
12
  durableStream?: string;
13
+ /** Receives chunks tools write via `toolWriter()`; those that are part of the message are checkpointed and re-emitted on replay, before the tool's output. */
14
+ writer?: UIMessageStreamWriter;
15
+ }
16
+
17
+ const currentWriter = new AsyncLocalStorage<UIMessageStreamWriter>();
18
+
19
+ /** Returns a UI message stream writer bound to the current `durableTools` tool call. */
20
+ export function toolWriter(): UIMessageStreamWriter {
21
+ const writer = currentWriter.getStore();
22
+ if (!writer) throw new Error('toolWriter() can only be called from a tool wrapped by durableTools.');
23
+ return writer;
24
+ }
25
+
26
+ // Tool output plus the message chunks its call wrote, as checkpointed; a bare output is a call that wrote none.
27
+ interface ToolEnvelope {
28
+ __dbosToolChunks: 1;
29
+ output: unknown;
30
+ chunks: UIMessageChunk[];
31
+ }
32
+
33
+ function isToolEnvelope(value: unknown): value is ToolEnvelope {
34
+ return typeof value === 'object' && value !== null && (value as Partial<ToolEnvelope>).__dbosToolChunks === 1;
35
+ }
36
+
37
+ function isTransient(chunk: UIMessageChunk): boolean {
38
+ return (chunk as { transient?: boolean }).transient === true;
39
+ }
40
+
41
+ // An async generator's body runs as it is iterated, outside run(); bind each step of the iteration to the writer.
42
+ function runWithWriter(writer: UIMessageStreamWriter, fn: () => unknown): unknown {
43
+ const result = currentWriter.run(writer, fn);
44
+ if (!isAsyncIterable(result)) return result;
45
+ const iterator = currentWriter.run(writer, () => result[Symbol.asyncIterator]());
46
+ return {
47
+ [Symbol.asyncIterator]() {
48
+ return this;
49
+ },
50
+ next: (...args: [] | [unknown]) => currentWriter.run(writer, () => iterator.next(...args)),
51
+ return: (value?: unknown) => currentWriter.run(writer, () => iterator.return?.(value) ?? Promise.resolve({ done: true as const, value })),
52
+ throw: (error?: unknown) => currentWriter.run(writer, () => iterator.throw?.(error) ?? Promise.reject(error)),
53
+ };
54
+ }
55
+
56
+ // Outside a step: every chunk goes straight to the writer, or nowhere.
57
+ function liveWriter(writer: UIMessageStreamWriter | undefined): UIMessageStreamWriter {
58
+ return writer ?? { write() {}, merge() {}, onError: undefined };
59
+ }
60
+
61
+ /** One attempt's writer: transient chunks go out live, the rest wait for the step to succeed. */
62
+ class StepWriter implements UIMessageStreamWriter {
63
+ readonly chunks: UIMessageChunk[] = [];
64
+ private pending: Promise<unknown>[] = [];
65
+
66
+ constructor(
67
+ private readonly writer: UIMessageStreamWriter | undefined,
68
+ private readonly durableStream: string | undefined,
69
+ ) {}
70
+
71
+ get onError() {
72
+ return this.writer?.onError;
73
+ }
74
+
75
+ write(chunk: UIMessageChunk): void {
76
+ if (!isTransient(chunk)) {
77
+ this.chunks.push(chunk);
78
+ return;
79
+ }
80
+ this.writer?.write(chunk);
81
+ if (this.durableStream) this.track(writeDurableStream(this.durableStream, [chunk]));
82
+ }
83
+
84
+ merge(stream: ReadableStream<UIMessageChunk>): void {
85
+ this.track(
86
+ (async () => {
87
+ for await (const chunk of stream) this.write(chunk);
88
+ })(),
89
+ );
90
+ }
91
+
92
+ // Observed now so a failure after the tool throws is not an unhandled rejection; settle still sees it.
93
+ private track(promise: Promise<unknown>): void {
94
+ promise.catch(() => {});
95
+ this.pending.push(promise);
96
+ }
97
+
98
+ /** Waits for merges and live writes, including any started while waiting; a failed one fails the call. */
99
+ async settle(): Promise<void> {
100
+ while (this.pending.length > 0) {
101
+ const batch = this.pending;
102
+ this.pending = [];
103
+ await Promise.all(batch);
104
+ }
105
+ }
12
106
  }
13
107
 
14
108
  // Loose view of a tool's execute; the AI SDK validates input and supplies the options.
@@ -20,7 +114,8 @@ type ToolExecute = (input: unknown, options: { toolCallId: string; abortSignal?:
20
114
  * Retries are off by default (the AI SDK never retries tools); opt in per tool with `retriesAllowed`.
21
115
  */
22
116
  export function durableTools<TOOLS extends ToolSet>(tools: TOOLS, options: DurableToolsOptions = {}): TOOLS {
23
- const { tools: perTool, durableStream, ...defaults } = options;
117
+ const { tools: perTool, durableStream, writer, ...defaults } = options;
118
+ const live = liveWriter(writer);
24
119
  const durable: ToolSet = {};
25
120
  for (const [name, definition] of Object.entries(tools)) {
26
121
  const override = perTool?.[name];
@@ -30,11 +125,18 @@ export function durableTools<TOOLS extends ToolSet>(tools: TOOLS, options: Durab
30
125
  durable[name] = durableStream ? bindAgentTool(durableStream) : definition;
31
126
  continue;
32
127
  }
33
- if (typeof definition.execute !== 'function' || override === false) {
128
+ if (typeof definition.execute !== 'function') {
34
129
  durable[name] = definition;
35
130
  continue;
36
131
  }
37
132
  const execute = definition.execute as ToolExecute;
133
+ if (override === false) {
134
+ durable[name] = {
135
+ ...definition,
136
+ execute: (input: unknown, execOptions: Parameters<ToolExecute>[1]) => runWithWriter(live, () => execute(input, execOptions)),
137
+ } as ToolSet[string];
138
+ continue;
139
+ }
38
140
  const merged: StepConfig = { ...defaults, ...override };
39
141
  // Default classification (aborts and provider-declared non-retryable errors are terminal), but retries stay opt-in.
40
142
  const stepConfig: StepConfig = { ...withErrorClassification(merged), retriesAllowed: merged.retriesAllowed ?? false };
@@ -43,7 +145,7 @@ export function durableTools<TOOLS extends ToolSet>(tools: TOOLS, options: Durab
43
145
  ...definition,
44
146
  execute: (input: unknown, execOptions: Parameters<ToolExecute>[1]) => {
45
147
  assertNotInTransaction(name);
46
- if (!isInWorkflowFunction()) return execute(input, execOptions);
148
+ if (!isInWorkflowFunction()) return runWithWriter(live, () => execute(input, execOptions));
47
149
  const signal = execOptions.abortSignal;
48
150
  // An aborted call is done whatever the failure looks like; a retry would re-run a cancelled side effect.
49
151
  const callConfig: StepConfig = {
@@ -54,27 +156,43 @@ export function durableTools<TOOLS extends ToolSet>(tools: TOOLS, options: Durab
54
156
  return runDurableStep(
55
157
  `${prefix}.${execOptions.toolCallId}`,
56
158
  async () => {
57
- // A timed-out attempt is abandoned by DBOS but keeps running; forward its signal so the tool stops too.
58
- const timeoutSignal = DBOS.stepStatus?.timeoutSignal;
59
- const abortSignal = timeoutSignal && signal ? AbortSignal.any([signal, timeoutSignal]) : (timeoutSignal ?? signal);
159
+ // A timed-out attempt is abandoned by DBOS but keeps running, and a cancelled workflow's tool should stop: forward both signals.
160
+ const cancelSignal = stepCancelSignal();
161
+ const abortSignal = anySignal(signal, DBOS.stepStatus?.timeoutSignal, cancelSignal);
162
+ const stepWriter = new StepWriter(writer, durableStream);
163
+ const record = async (outcome: { output: unknown } | { errorText: string }, chunks?: UIMessageChunk[]) => {
164
+ if (!durableStream) return;
165
+ const timeout = await timedOutOutcome(callConfig.shouldRetry);
166
+ if (timeout === null) return;
167
+ await writeToolRecord(durableStream, execOptions.toolCallId, timeout ?? outcome, timeout ? [] : chunks);
168
+ };
60
169
  let output: unknown;
61
170
  try {
62
- output = await execute(input, abortSignal === signal ? execOptions : { ...execOptions, abortSignal });
63
- // A streaming execute can't checkpoint mid-flight; drain it and record the final value (the last yield).
64
- if (isAsyncIterable(output)) {
171
+ output = await currentWriter.run(stepWriter, async () => {
172
+ const value = await execute(input, abortSignal === signal ? execOptions : { ...execOptions, abortSignal });
173
+ // A streaming execute can't checkpoint mid-flight; drain it and record the final value (the last yield).
174
+ if (!isAsyncIterable(value)) return value;
65
175
  let last: unknown;
66
- for await (last of output);
67
- output = last;
68
- }
176
+ for await (last of value);
177
+ return last;
178
+ });
179
+ await stepWriter.settle();
69
180
  } catch (error) {
70
- if (durableStream) await writeToolRecord(durableStream, execOptions.toolCallId, { errorText: errorMessage(error) });
181
+ // A cancelled workflow's reader ends the turn with an abort; this call has no outcome of its own.
182
+ if (!cancelSignal?.aborted) await record({ errorText: errorMessage(error) });
71
183
  throw error;
72
184
  }
73
- if (durableStream) await writeToolRecord(durableStream, execOptions.toolCallId, { output });
74
- return output;
185
+ const chunks = stepWriter.chunks;
186
+ await record({ output }, chunks);
187
+ return chunks.length > 0 ? ({ __dbosToolChunks: 1, output, chunks } satisfies ToolEnvelope) : output;
75
188
  },
76
189
  callConfig,
77
- );
190
+ ).then((result) => {
191
+ if (!isToolEnvelope(result)) return result;
192
+ // Runs on first execution and on replay alike, so the workflow's message gets the same parts either way.
193
+ for (const chunk of result.chunks) writer?.write(chunk);
194
+ return result.output;
195
+ });
78
196
  },
79
197
  } as ToolSet[string];
80
198
  }