@dbos-inc/vercel-ai 0.1.5 → 0.3.7

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.
package/src/middleware.ts CHANGED
@@ -21,21 +21,26 @@ import type {
21
21
  SharedV4Warning,
22
22
  } from '@ai-sdk/provider' with { 'resolution-mode': 'import' };
23
23
  import { assertNotInTransaction, isInWorkflowFunction, restoreAISDKErrorIdentity, withErrorClassification } from './internal';
24
+ import { type DurableStreamOptions, ModelStreamWriter, resolveDurableStream } from './durable-stream';
25
+
26
+ export interface DurableCallsOptions extends StepConfig {
27
+ /** Write each streamed model call's parts to this durable stream from inside its step (see readDurableStream). */
28
+ durableStream?: DurableStreamOptions;
29
+ }
24
30
 
25
31
  // In-flight durable model calls per workflow; concurrent calls have a nondeterministic DBOS step order on replay, so we reject them.
26
32
  const inflightModelCalls = new Map<string, number>();
33
+ // A stream call whose consumer detached (abort/cancel) but whose step is still settling; the next call in that workflow waits for it.
34
+ const settlingModelCalls = new Map<string, Promise<void>>();
27
35
 
28
- function enterDurableModelCall(operation: 'generate' | 'stream' | 'embed'): string {
36
+ async function enterDurableModelCall(): Promise<string> {
29
37
  const workflowID = DBOS.workflowID!;
38
+ // Sequence after a detached call's checkpoint, so steps are always recorded in the order they were started.
39
+ await settlingModelCalls.get(workflowID);
30
40
  const inflight = inflightModelCalls.get(workflowID) ?? 0;
31
41
  if (inflight > 0) {
32
- // embedMany parallelizes its batches; maxParallelCalls: 1 serializes them deterministically. Other callers use child workflows.
33
- const remedy =
34
- operation === 'embed'
35
- ? 'pass maxParallelCalls: 1 to embedMany, or run each call in its own child workflow with DBOS.startWorkflow'
36
- : 'run each call in its own child workflow with DBOS.startWorkflow';
37
42
  throw new Error(
38
- `Concurrent durable model calls in workflow "${workflowID}" are not supported because their step order is nondeterministic on replay; ${remedy}.`,
43
+ `Concurrent durable model calls in workflow "${workflowID}" are not supported because their step order is nondeterministic on replay; run each call in its own child workflow with DBOS.startWorkflow.`,
39
44
  );
40
45
  }
41
46
  inflightModelCalls.set(workflowID, inflight + 1);
@@ -52,21 +57,36 @@ function exitDurableModelCall(workflowID: string): void {
52
57
  }
53
58
 
54
59
  /** AI SDK language-model middleware that runs each model call as a durable, checkpointed DBOS step (replayed on recovery); outside a workflow it calls the model directly. */
55
- export function durableCalls(options: StepConfig = {}): LanguageModelMiddleware {
56
- const stepConfig = withErrorClassification(options);
60
+ export function durableCalls(options: DurableCallsOptions = {}): LanguageModelMiddleware {
61
+ const { durableStream, ...stepOptions } = options;
62
+ const stepConfig = withErrorClassification(stepOptions);
63
+ const streamConfig = resolveDurableStream(durableStream);
57
64
  return {
58
65
  specificationVersion: 'v4',
59
66
 
60
- wrapGenerate: async ({ doGenerate, model }) => {
67
+ wrapGenerate: async ({ doGenerate, params, model }) => {
61
68
  assertNotInTransaction('generate');
62
69
  if (!isInWorkflowFunction()) {
63
70
  return await doGenerate();
64
71
  }
65
- const workflowID = enterDurableModelCall('generate');
72
+ const workflowID = await enterDurableModelCall();
66
73
  try {
67
- return await DBOS.runStep(async () => ensureResponseMetadata(encodeBinaryContent(await doGenerate())), {
74
+ return await DBOS.runStep(
75
+ async () => {
76
+ const result = ensureResponseMetadata(encodeBinaryContent(await doGenerate()));
77
+ // A non-streaming call writes its whole output at once, so the stream holds every call the loop makes, not only streamed ones.
78
+ if (streamConfig) {
79
+ const streamWriter = new ModelStreamWriter(streamConfig);
80
+ for (const part of replayParts(result)) streamWriter.push(part);
81
+ await streamWriter.end({ finishReason: result.finishReason });
82
+ }
83
+ return result;
84
+ },
85
+ {
68
86
  ...stepConfig,
69
87
  name: stepConfig.name ?? stepName(model, 'generate'),
88
+ // An aborted call is never retried, whatever the provider's rejection looks like.
89
+ shouldRetry: async (error: unknown) => params.abortSignal?.aborted !== true && (await stepConfig.shouldRetry!(error)),
70
90
  });
71
91
  } catch (error) {
72
92
  // Restore the AI SDK error identity a replay revival strips, so the SDK's retry/catch logic behaves the same.
@@ -81,39 +101,36 @@ export function durableCalls(options: StepConfig = {}): LanguageModelMiddleware
81
101
  if (!isInWorkflowFunction()) {
82
102
  return await doStream();
83
103
  }
84
- const workflowID = enterDurableModelCall('stream');
85
- // An aborted consumer is done with this call, like a cancelled one: post-abort failures must checkpoint as a (partial) success, or replay would fail where the live run ended gracefully.
86
- const aborted = () => params.abortSignal?.aborted === true;
87
-
88
- // Free the concurrency guard as soon as the consumer detaches (abort/cancel), not only when the step
89
- // settles: this call's step is already sequenced, so a sequential follow-up in the same workflow is
90
- // deterministic on replay and must not be rejected as concurrent while this step drains in the background.
91
- let guardReleased = false;
92
- const releaseGuard = () => {
93
- if (guardReleased) return;
94
- guardReleased = true;
95
- params.abortSignal?.removeEventListener('abort', releaseGuard);
96
- exitDurableModelCall(workflowID);
97
- };
98
- params.abortSignal?.addEventListener('abort', releaseGuard, { once: true });
104
+ const workflowID = await enterDurableModelCall();
105
+ // The caller's abortSignal merged with the AI SDK's timeouts; an abort stops the provider call and is recorded as the step's outcome.
106
+ const abortSignal = params.abortSignal;
107
+ const aborted = () => abortSignal?.aborted === true;
99
108
 
100
109
  let executed = false;
101
110
  let cancelled = false;
102
111
  let emittedLive = false;
112
+ let released = false;
103
113
  let controller!: ReadableStreamDefaultController<LanguageModelV4StreamPart>;
104
114
  let step!: Promise<LanguageModelV4GenerateResult>;
115
+ let settled!: Promise<void>;
116
+ // A detached consumer (abort/cancel) may issue a sequential follow-up while this step still settles: make that call wait for this checkpoint instead of refusing it as concurrent.
117
+ const detach = () => {
118
+ abortSignal?.removeEventListener('abort', detach);
119
+ if (!released) settlingModelCalls.set(workflowID, settled);
120
+ };
121
+ abortSignal?.addEventListener('abort', detach, { once: true });
105
122
  const stream = new ReadableStream<LanguageModelV4StreamPart>({
106
123
  start(c) {
107
124
  controller = c;
108
125
  },
109
- // Await the step so an early cancel still blocks until the model result is checkpointed.
126
+ // Await the step so a direct cancel blocks until the checkpoint; through the AI SDK's pipeline a consumer's early exit returns sooner, hence the rule to drain or abort instead.
110
127
  async cancel() {
111
128
  cancelled = true;
112
- releaseGuard();
129
+ detach();
113
130
  // Post-cancel failures normally checkpoint as a success; anything else (e.g. a failed checkpoint write) is only visible here.
114
- await step.catch((error: unknown) =>
115
- DBOS.logger.warn(`Durable model call step failed after consumer cancel: ${String(error)}`),
116
- );
131
+ await step.catch((error: unknown) => {
132
+ if (!aborted()) DBOS.logger.warn(`Durable model call step failed after consumer cancel: ${String(error)}`);
133
+ });
117
134
  },
118
135
  });
119
136
  const emit = (part: LanguageModelV4StreamPart) => {
@@ -124,11 +141,11 @@ export function durableCalls(options: StepConfig = {}): LanguageModelMiddleware
124
141
  }
125
142
  };
126
143
 
127
- // Once any output part has streamed live, a retry would re-stream from scratch and duplicate output, so stop retrying.
144
+ // Once any output part has streamed live, a retry would re-stream from scratch and duplicate output; an aborted call is never retried.
128
145
  const streamStepConfig: StepConfig = {
129
146
  ...stepConfig,
130
147
  shouldRetry: async (error: unknown) =>
131
- !emittedLive && (stepConfig.shouldRetry ? await stepConfig.shouldRetry(error) : true),
148
+ !emittedLive && !aborted() && (stepConfig.shouldRetry ? await stepConfig.shouldRetry(error) : true),
132
149
  };
133
150
 
134
151
  try {
@@ -143,23 +160,26 @@ export function durableCalls(options: StepConfig = {}): LanguageModelMiddleware
143
160
  let sawFinish = false;
144
161
  const abandon = () => void reader?.cancel().catch(() => {});
145
162
  timeoutSignal?.addEventListener('abort', abandon, { once: true });
146
- // A consumer abort detaches this call; stop draining now so the checkpoint lands promptly and matches what streamed.
147
- params.abortSignal?.addEventListener('abort', abandon, { once: true });
163
+ // A consumer abort tears the provider call down too; the attempt is then recorded as aborted below.
164
+ abortSignal?.addEventListener('abort', abandon, { once: true });
165
+ // Step-scope stream writes: cheap, replay-safe, and never duplicated since a retry is refused once content has streamed.
166
+ const streamWriter = streamConfig ? new ModelStreamWriter(streamConfig) : undefined;
148
167
  try {
149
168
  streamResult = await doStream();
150
169
  reader = streamResult.stream.getReader();
151
170
  for (;;) {
152
171
  const { done, value: part } = await reader.read();
153
172
  if (timeoutSignal?.aborted) throw (timeoutSignal.reason ?? new Error('step attempt timed out'));
154
- if (done) break;
173
+ if (done || aborted()) break;
155
174
  if (part.type === 'error') {
156
- // A cancelled or aborted consumer abandoned this call; don't let a late failure become the step outcome, or replay would fail where the live run succeeded.
157
- if (cancelled || aborted()) break;
175
+ // 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
+ if (cancelled) break;
158
177
  throw toStepError(part.error);
159
178
  }
160
179
  // Stream deltas live but withhold 'finish' until the checkpoint is durable: the AI SDK runs tool calls (and their durable steps) on 'finish', which must not checkpoint before this model step.
161
180
  if (part.type === 'finish') sawFinish = true;
162
181
  else emit(part);
182
+ streamWriter?.push(part);
163
183
  accumulator.add(part);
164
184
  }
165
185
  // No terminal part and no output: fail (retryably) like the AI SDK's NoOutputGeneratedError, instead of checkpointing a permanent empty success.
@@ -167,34 +187,55 @@ export function durableCalls(options: StepConfig = {}): LanguageModelMiddleware
167
187
  throw new Error('Model stream ended without a finish part or any output.');
168
188
  }
169
189
  } catch (error) {
170
- // Same rule for stream-level failures (doStream or a read rejecting) after a cancel or abort.
171
- if (!cancelled && !aborted()) throw 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()) {
192
+ await streamWriter?.abandon();
193
+ throw error;
194
+ }
172
195
  } finally {
173
196
  timeoutSignal?.removeEventListener('abort', abandon);
174
- params.abortSignal?.removeEventListener('abort', abandon);
197
+ abortSignal?.removeEventListener('abort', abandon);
175
198
  // Tear down the provider stream on early exits (error part, post-cancel break); a no-op after a clean drain.
176
199
  void reader?.cancel().catch(() => {});
177
200
  }
201
+ // Record the abort as the step's failure; replay rethrows it, so the workflow must catch aborts it means to survive.
202
+ if (aborted()) {
203
+ await streamWriter?.end({ aborted: true });
204
+ throw toAbortError(abortSignal!.reason);
205
+ }
178
206
  // Give the response a durable id/timestamp when the provider sent none, and emit it live (before the
179
207
  // withheld 'finish') so the SDK sees the same values live and on replay instead of a fresh fallback.
180
208
  // Skip the live emit for a timed-out (abandoned) attempt so it can't interleave with its retry.
181
209
  const responseMetadataPart = accumulator.fillResponseMetadata(randomUUID(), new Date());
182
210
  if (responseMetadataPart && !timeoutSignal?.aborted) emit(responseMetadataPart);
183
- return encodeBinaryContent(accumulator.result(streamResult?.request, streamResult?.response));
211
+ const recorded = encodeBinaryContent(accumulator.result(streamResult?.request, streamResult?.response));
212
+ // Every stream write lands before the checkpoint, so a reader that sees the next step has seen all of this one.
213
+ await streamWriter?.end({ finishReason: recorded.finishReason });
214
+ return recorded;
184
215
  },
185
216
  { ...streamStepConfig, name: stepConfig.name ?? stepName(model, 'stream') },
186
217
  );
187
218
  } catch (error) {
188
219
  // runStep can throw synchronously (e.g. a shutdown race); don't leak the guard entry.
189
- releaseGuard();
220
+ abortSignal?.removeEventListener('abort', detach);
221
+ exitDurableModelCall(workflowID);
190
222
  throw error;
191
223
  }
224
+ // Hold the guard until the step has settled, so the next call in this workflow always starts after this checkpoint.
225
+ settled = step.then(
226
+ () => undefined,
227
+ () => undefined,
228
+ ).then(() => {
229
+ released = true;
230
+ abortSignal?.removeEventListener('abort', detach);
231
+ exitDurableModelCall(workflowID);
232
+ if (settlingModelCalls.get(workflowID) === settled) settlingModelCalls.delete(workflowID);
233
+ });
192
234
 
193
235
  // Drive the returned stream from the settled step: a live run emits only the withheld 'finish' (deltas already
194
236
  // streamed); a recovered run synthesizes the whole stream from the checkpoint. Either way consumers finish
195
237
  // only after the result is durable.
196
- void step
197
- .then(
238
+ void step.then(
198
239
  (recorded) => {
199
240
  if (cancelled) return;
200
241
  if (executed) {
@@ -210,10 +251,11 @@ export function durableCalls(options: StepConfig = {}): LanguageModelMiddleware
210
251
  controller.close();
211
252
  },
212
253
  (error: unknown) => {
213
- if (!cancelled) controller.error(restoreAISDKErrorIdentity(error));
254
+ if (cancelled) return;
255
+ // Live abort: surface the signal's own reason so the AI SDK takes its abort path; a replayed abort is an ordinary error.
256
+ controller.error(executed && aborted() ? (abortSignal?.reason ?? error) : restoreAISDKErrorIdentity(error));
214
257
  },
215
- )
216
- .finally(releaseGuard);
258
+ );
217
259
 
218
260
  return { stream };
219
261
  },
@@ -225,12 +267,17 @@ export function durableEmbeddingCalls(options: StepConfig = {}): EmbeddingModelM
225
267
  const stepConfig = withErrorClassification(options);
226
268
  return {
227
269
  specificationVersion: 'v4',
270
+ // embedMany awaits this per call: inside a workflow the batches run sequentially (deterministic step order on replay); elsewhere the model's own answer stands.
271
+ overrideSupportsParallelCalls: ({ model }) => ({
272
+ then: (onfulfilled, onrejected) =>
273
+ Promise.resolve(isInWorkflowFunction() ? false : model.supportsParallelCalls).then(onfulfilled, onrejected),
274
+ }),
228
275
  wrapEmbed: async ({ doEmbed, model }) => {
229
276
  assertNotInTransaction('embed');
230
277
  if (!isInWorkflowFunction()) {
231
278
  return await doEmbed();
232
279
  }
233
- const workflowID = enterDurableModelCall('embed');
280
+ const workflowID = await enterDurableModelCall();
234
281
  try {
235
282
  return await DBOS.runStep(async () => doEmbed(), {
236
283
  ...stepConfig,
@@ -271,6 +318,12 @@ export function durableImageCalls(options: StepConfig = {}): ImageModelMiddlewar
271
318
  };
272
319
  }
273
320
 
321
+ // The abort reason as a recordable Error (an AbortSignal's reason may be any value); the name is what the AI SDK's abort checks read.
322
+ function toAbortError(reason: unknown): Error {
323
+ if (reason instanceof Error) return reason;
324
+ return Object.assign(new Error(String(reason)), { name: 'AbortError' });
325
+ }
326
+
274
327
  /** Convert generated image bytes (Uint8Array) to base64 (spec-allowed) to keep checkpoints compact. */
275
328
  function encodeImageResult(result: ImageModelV4Result): ImageModelV4Result {
276
329
  const images = result.images as (string | Uint8Array)[];
package/src/tools.ts ADDED
@@ -0,0 +1,86 @@
1
+ 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';
4
+ import { AGENT_TOOL } from './agent-tool';
5
+ import { writeToolRecord } from './durable-stream';
6
+
7
+ export interface DurableToolsOptions extends StepConfig {
8
+ /** Per-tool step config overriding the defaults; `false` leaves that tool non-durable. */
9
+ tools?: Record<string, StepConfig | false>;
10
+ /** Write each tool call's output (or error) to this durable stream from inside its step. */
11
+ durableStream?: string;
12
+ }
13
+
14
+ // Loose view of a tool's execute; the AI SDK validates input and supplies the options.
15
+ type ToolExecute = (input: unknown, options: { toolCallId: string; abortSignal?: AbortSignal }) => unknown;
16
+
17
+ /**
18
+ * Wraps each tool's `execute` so that, inside a workflow, every tool call runs as a durable DBOS step named
19
+ * `<tool>.<toolCallId>` and replays from its checkpoint on recovery; outside a workflow tools run unchanged.
20
+ * Retries are off by default (the AI SDK never retries tools); opt in per tool with `retriesAllowed`.
21
+ */
22
+ export function durableTools<TOOLS extends ToolSet>(tools: TOOLS, options: DurableToolsOptions = {}): TOOLS {
23
+ const { tools: perTool, durableStream, ...defaults } = options;
24
+ const durable: ToolSet = {};
25
+ for (const [name, definition] of Object.entries(tools)) {
26
+ const override = perTool?.[name];
27
+ // An agent tool is a child workflow, not a step: leave it unwrapped, binding this durable stream to it.
28
+ const bindAgentTool = (definition as { [AGENT_TOOL]?: (key: string) => ToolSet[string] })[AGENT_TOOL];
29
+ if (bindAgentTool) {
30
+ durable[name] = durableStream ? bindAgentTool(durableStream) : definition;
31
+ continue;
32
+ }
33
+ if (typeof definition.execute !== 'function' || override === false) {
34
+ durable[name] = definition;
35
+ continue;
36
+ }
37
+ const execute = definition.execute as ToolExecute;
38
+ const merged: StepConfig = { ...defaults, ...override };
39
+ // Default classification (aborts and provider-declared non-retryable errors are terminal), but retries stay opt-in.
40
+ const stepConfig: StepConfig = { ...withErrorClassification(merged), retriesAllowed: merged.retriesAllowed ?? false };
41
+ const prefix = stepConfig.name ?? name;
42
+ durable[name] = {
43
+ ...definition,
44
+ execute: (input: unknown, execOptions: Parameters<ToolExecute>[1]) => {
45
+ assertNotInTransaction(name);
46
+ if (!isInWorkflowFunction()) return execute(input, execOptions);
47
+ const signal = execOptions.abortSignal;
48
+ // An aborted call is done whatever the failure looks like; a retry would re-run a cancelled side effect.
49
+ const callConfig: StepConfig = {
50
+ ...stepConfig,
51
+ shouldRetry: async (error: unknown) => !signal?.aborted && (await stepConfig.shouldRetry!(error)),
52
+ };
53
+ // The tool call id comes from the checkpointed model result, so a reordered parallel step fails replay instead of swapping results.
54
+ return runDurableStep(
55
+ `${prefix}.${execOptions.toolCallId}`,
56
+ 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);
60
+ let output: unknown;
61
+ 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)) {
65
+ let last: unknown;
66
+ for await (last of output);
67
+ output = last;
68
+ }
69
+ } catch (error) {
70
+ if (durableStream) await writeToolRecord(durableStream, execOptions.toolCallId, { errorText: errorMessage(error) });
71
+ throw error;
72
+ }
73
+ if (durableStream) await writeToolRecord(durableStream, execOptions.toolCallId, { output });
74
+ return output;
75
+ },
76
+ callConfig,
77
+ );
78
+ },
79
+ } as ToolSet[string];
80
+ }
81
+ return durable as TOOLS;
82
+ }
83
+
84
+ function errorMessage(error: unknown): string {
85
+ return error instanceof Error ? error.message : String(error);
86
+ }