@tanstack/ai 0.43.1 → 0.44.0

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.
Files changed (69) hide show
  1. package/dist/esm/activities/chat/messages.js +21 -8
  2. package/dist/esm/activities/chat/messages.js.map +1 -1
  3. package/dist/esm/activities/embed/adapter.d.ts +69 -0
  4. package/dist/esm/activities/embed/adapter.js +23 -0
  5. package/dist/esm/activities/embed/adapter.js.map +1 -0
  6. package/dist/esm/activities/embed/index.d.ts +117 -0
  7. package/dist/esm/activities/embed/index.js +166 -0
  8. package/dist/esm/activities/embed/index.js.map +1 -0
  9. package/dist/esm/activities/error-payload.d.ts +8 -0
  10. package/dist/esm/activities/error-payload.js +29 -17
  11. package/dist/esm/activities/error-payload.js.map +1 -1
  12. package/dist/esm/activities/generateAudio/index.d.ts +12 -0
  13. package/dist/esm/activities/generateAudio/index.js +19 -6
  14. package/dist/esm/activities/generateAudio/index.js.map +1 -1
  15. package/dist/esm/activities/generateImage/index.d.ts +12 -0
  16. package/dist/esm/activities/generateImage/index.js +21 -7
  17. package/dist/esm/activities/generateImage/index.js.map +1 -1
  18. package/dist/esm/activities/generateSpeech/index.d.ts +17 -1
  19. package/dist/esm/activities/generateSpeech/index.js +19 -6
  20. package/dist/esm/activities/generateSpeech/index.js.map +1 -1
  21. package/dist/esm/activities/generateTranscription/index.d.ts +17 -1
  22. package/dist/esm/activities/generateTranscription/index.js +19 -6
  23. package/dist/esm/activities/generateTranscription/index.js.map +1 -1
  24. package/dist/esm/activities/generateVideo/index.d.ts +18 -0
  25. package/dist/esm/activities/generateVideo/index.js +54 -15
  26. package/dist/esm/activities/generateVideo/index.js.map +1 -1
  27. package/dist/esm/activities/index.d.ts +8 -2
  28. package/dist/esm/activities/index.js +11 -7
  29. package/dist/esm/activities/middleware/types.d.ts +1 -1
  30. package/dist/esm/activities/rerank/adapter.d.ts +63 -0
  31. package/dist/esm/activities/rerank/adapter.js +23 -0
  32. package/dist/esm/activities/rerank/adapter.js.map +1 -0
  33. package/dist/esm/activities/rerank/index.d.ts +92 -0
  34. package/dist/esm/activities/rerank/index.js +163 -0
  35. package/dist/esm/activities/rerank/index.js.map +1 -0
  36. package/dist/esm/activities/summarize/index.d.ts +17 -1
  37. package/dist/esm/activities/summarize/index.js +19 -5
  38. package/dist/esm/activities/summarize/index.js.map +1 -1
  39. package/dist/esm/index.d.ts +7 -2
  40. package/dist/esm/index.js +5 -1
  41. package/dist/esm/middlewares/otel.js +20 -2
  42. package/dist/esm/middlewares/otel.js.map +1 -1
  43. package/dist/esm/types.d.ts +195 -0
  44. package/dist/esm/utilities/activity-abort.d.ts +53 -0
  45. package/dist/esm/utilities/activity-abort.js +150 -0
  46. package/dist/esm/utilities/activity-abort.js.map +1 -0
  47. package/dist/esm/utilities/embedding-input.d.ts +32 -0
  48. package/dist/esm/utilities/embedding-input.js +61 -0
  49. package/dist/esm/utilities/embedding-input.js.map +1 -0
  50. package/package.json +3 -3
  51. package/src/activities/chat/messages.ts +30 -1
  52. package/src/activities/embed/adapter.ts +112 -0
  53. package/src/activities/embed/index.ts +318 -0
  54. package/src/activities/error-payload.ts +41 -9
  55. package/src/activities/generateAudio/index.ts +47 -5
  56. package/src/activities/generateImage/index.ts +48 -5
  57. package/src/activities/generateSpeech/index.ts +52 -9
  58. package/src/activities/generateTranscription/index.ts +52 -9
  59. package/src/activities/generateVideo/index.ts +131 -33
  60. package/src/activities/index.ts +44 -0
  61. package/src/activities/middleware/types.ts +2 -0
  62. package/src/activities/rerank/adapter.ts +90 -0
  63. package/src/activities/rerank/index.ts +302 -0
  64. package/src/activities/summarize/index.ts +59 -19
  65. package/src/index.ts +19 -0
  66. package/src/middlewares/otel.ts +38 -3
  67. package/src/types.ts +219 -0
  68. package/src/utilities/activity-abort.ts +197 -0
  69. package/src/utilities/embedding-input.ts +83 -0
@@ -0,0 +1,318 @@
1
+ /**
2
+ * Embed Activity
3
+ *
4
+ * Generates embedding vectors from text and (for multimodal models) image
5
+ * inputs. This is a self-contained module with implementation, types, and JSDoc.
6
+ */
7
+
8
+ import { aiEventClient } from '@tanstack/ai-event-client'
9
+ import { resolveDebugOption } from '../../logger/resolve'
10
+ import {
11
+ createGenerationContext,
12
+ runGenerationError,
13
+ runGenerationFinish,
14
+ runGenerationStart,
15
+ runGenerationUsage,
16
+ } from '../middleware/run'
17
+ import { countEmbeddingInputModalities } from '../../utilities/embedding-input'
18
+ import type { InternalLogger } from '../../logger/internal-logger'
19
+ import type { DebugOption } from '../../logger/types'
20
+ import type { GenerationMiddleware } from '../middleware/types'
21
+ import type { EmbeddingAdapter } from './adapter'
22
+ import type {
23
+ EmbeddingInputItem,
24
+ EmbeddingInputItemFor,
25
+ EmbeddingResult,
26
+ } from '../../types'
27
+
28
+ // ===========================
29
+ // Activity Kind
30
+ // ===========================
31
+
32
+ /** The adapter kind this activity handles */
33
+ export const kind = 'embedding' as const
34
+
35
+ // ===========================
36
+ // Type Extraction Helpers
37
+ // ===========================
38
+
39
+ /**
40
+ * Extract model-specific provider options from an EmbeddingAdapter via ~types.
41
+ * If the model has specific options defined in ModelProviderOptions (and not just via index signature),
42
+ * use those; otherwise fall back to base provider options.
43
+ */
44
+ export type EmbedProviderOptionsForModel<TAdapter, TModel extends string> =
45
+ TAdapter extends EmbeddingAdapter<
46
+ any,
47
+ infer BaseOptions,
48
+ infer ModelOptions,
49
+ any
50
+ >
51
+ ? string extends keyof ModelOptions
52
+ ? // ModelOptions is Record<string, unknown> or has index signature - use BaseOptions
53
+ BaseOptions
54
+ : // ModelOptions has explicit keys - check if TModel is one of them
55
+ TModel extends keyof ModelOptions
56
+ ? ModelOptions[TModel]
57
+ : BaseOptions
58
+ : object
59
+
60
+ /**
61
+ * Extract the input type a model accepts from an EmbeddingAdapter via ~types.
62
+ * Adapters declare a per-model input-modality map; models in the map get an
63
+ * `input` narrowed to their supported item types (text-only models accept
64
+ * `string | TextPart`), so unsupported items fail at compile time. Adapters
65
+ * without a map fall back to the full EmbeddingInputItem union.
66
+ */
67
+ export type EmbeddingInputForModel<TAdapter, TModel extends string> =
68
+ TAdapter extends EmbeddingAdapter<any, any, any, infer ModsByName>
69
+ ? string extends keyof ModsByName
70
+ ? // No explicit map - accept the full union
71
+ EmbeddingInputItem | Array<EmbeddingInputItem>
72
+ : TModel extends keyof ModsByName
73
+ ?
74
+ | EmbeddingInputItemFor<ModsByName[TModel][number]>
75
+ | Array<EmbeddingInputItemFor<ModsByName[TModel][number]>>
76
+ : EmbeddingInputItem | Array<EmbeddingInputItem>
77
+ : EmbeddingInputItem | Array<EmbeddingInputItem>
78
+
79
+ // ===========================
80
+ // Activity Options Type
81
+ // ===========================
82
+
83
+ /**
84
+ * Options for the embed activity.
85
+ * The model is extracted from the adapter's model property.
86
+ *
87
+ * @template TAdapter - The embedding adapter type
88
+ */
89
+ export type EmbedOptions<
90
+ TAdapter extends EmbeddingAdapter<string, any, any, any>,
91
+ > = {
92
+ /** The embedding adapter to use (must be created with a model) */
93
+ adapter: TAdapter & { kind: typeof kind }
94
+ /**
95
+ * What to embed: a single item or an array of items. Each item in the array
96
+ * produces exactly one vector. An item is a plain string, a text part, an
97
+ * image part, or — for models that embed text and image together — a fused
98
+ * item written as a nested array of parts (`[textPart, imagePart]`), the
99
+ * same `Array<ContentPart>` shape chat messages use. The accepted item types
100
+ * are narrowed per model via the adapter's input-modality map.
101
+ */
102
+ input: EmbeddingInputForModel<TAdapter, TAdapter['model']>
103
+ /**
104
+ * Requested output dimensionality. Supported by models with Matryoshka /
105
+ * configurable dimensions; adapters for fixed-dimension models throw a
106
+ * clear runtime error when this is set.
107
+ */
108
+ dimensions?: number
109
+ /**
110
+ * Enable debug logging. Pass `true` to enable all categories, `false` to
111
+ * silence everything including errors, or a `DebugConfig` object for granular
112
+ * control and/or a custom `Logger`.
113
+ */
114
+ debug?: DebugOption
115
+ /**
116
+ * Observe-only middleware notified on start, usage, success, and error. Pass
117
+ * `otelMiddleware()` to emit OpenTelemetry spans, or implement the
118
+ * `GenerationMiddleware` contract for a custom backend.
119
+ */
120
+ middleware?: Array<GenerationMiddleware>
121
+ } & ({} extends EmbedProviderOptionsForModel<TAdapter, TAdapter['model']>
122
+ ? {
123
+ /** Provider-specific options for embedding generation */ modelOptions?: EmbedProviderOptionsForModel<
124
+ TAdapter,
125
+ TAdapter['model']
126
+ >
127
+ }
128
+ : {
129
+ /** Provider-specific options for embedding generation */ modelOptions: EmbedProviderOptionsForModel<
130
+ TAdapter,
131
+ TAdapter['model']
132
+ >
133
+ })
134
+
135
+ function createId(prefix: string): string {
136
+ return `${prefix}-${Date.now()}-${Math.random().toString(36).slice(2, 9)}`
137
+ }
138
+
139
+ // ===========================
140
+ // Activity Implementation
141
+ // ===========================
142
+
143
+ /**
144
+ * Embed activity - generates embedding vectors from text and image inputs.
145
+ *
146
+ * Accepts a single item or an array of items; the result always carries an
147
+ * `embeddings` array with one vector per input item, in input order.
148
+ *
149
+ * @example Embed a single text
150
+ * ```ts
151
+ * import { embed } from '@tanstack/ai'
152
+ * import { openaiEmbedding } from '@tanstack/ai-openai'
153
+ *
154
+ * const result = await embed({
155
+ * adapter: openaiEmbedding('text-embedding-3-small'),
156
+ * input: 'a red guitar',
157
+ * })
158
+ *
159
+ * console.log(result.embeddings[0].vector)
160
+ * ```
161
+ *
162
+ * @example Batch with requested dimensions
163
+ * ```ts
164
+ * const result = await embed({
165
+ * adapter: openaiEmbedding('text-embedding-3-large'),
166
+ * input: ['a red guitar', 'a blue drum kit'],
167
+ * dimensions: 1024,
168
+ * })
169
+ * ```
170
+ *
171
+ * @example Multimodal embedding (text + image fused into one vector)
172
+ * ```ts
173
+ * import { cohereEmbedding } from '@tanstack/ai-cohere'
174
+ *
175
+ * // A nested array of parts fuses them into a single vector. The outer array
176
+ * // is the item list, so this embeds one fused item into one vector.
177
+ * const result = await embed({
178
+ * adapter: cohereEmbedding('embed-v4.0'),
179
+ * input: [
180
+ * [
181
+ * { type: 'text', content: 'product photo' },
182
+ * { type: 'image', source: { type: 'data', value: base64, mimeType: 'image/png' } },
183
+ * ],
184
+ * ],
185
+ * modelOptions: { inputType: 'search_document' },
186
+ * })
187
+ * ```
188
+ */
189
+ export async function embed<
190
+ TAdapter extends EmbeddingAdapter<string, any, any, any>,
191
+ >(options: EmbedOptions<TAdapter>): Promise<EmbeddingResult> {
192
+ const { adapter, middleware } = options
193
+ const model = adapter.model
194
+ const requestId = createId('embedding')
195
+ const startTime = Date.now()
196
+ const logger: InternalLogger = resolveDebugOption(options.debug)
197
+ const modelOptions = (options as { modelOptions?: Record<string, unknown> })
198
+ .modelOptions
199
+
200
+ // Normalize once: adapters always receive an array of items.
201
+ const inputItems: Array<EmbeddingInputItem> = Array.isArray(options.input)
202
+ ? options.input
203
+ : [options.input]
204
+ const { textInputCount, imageInputCount } =
205
+ countEmbeddingInputModalities(inputItems)
206
+
207
+ const mwCtx = createGenerationContext({
208
+ requestId,
209
+ activity: 'embedding',
210
+ provider: adapter.name,
211
+ model,
212
+ modelOptions,
213
+ createId,
214
+ })
215
+
216
+ await runGenerationStart(middleware, mwCtx)
217
+
218
+ aiEventClient.emit('embedding:request:started', {
219
+ requestId,
220
+ provider: adapter.name,
221
+ model,
222
+ inputCount: inputItems.length,
223
+ textInputCount,
224
+ imageInputCount,
225
+ dimensions: options.dimensions,
226
+ modelOptions,
227
+ timestamp: startTime,
228
+ })
229
+
230
+ logger.request(`activity=embed provider=${adapter.name} model=${model}`, {
231
+ provider: adapter.name,
232
+ model,
233
+ })
234
+
235
+ try {
236
+ const result = await adapter.createEmbeddings({
237
+ model,
238
+ input: inputItems,
239
+ dimensions: options.dimensions,
240
+ modelOptions,
241
+ logger,
242
+ })
243
+ const duration = Date.now() - startTime
244
+
245
+ aiEventClient.emit('embedding:request:completed', {
246
+ requestId,
247
+ provider: adapter.name,
248
+ model,
249
+ embeddingCount: result.embeddings.length,
250
+ dimensions: result.embeddings[0]?.vector.length,
251
+ duration,
252
+ modelOptions,
253
+ timestamp: Date.now(),
254
+ })
255
+
256
+ logger.output(`activity=embed count=${result.embeddings.length}`, {
257
+ embeddingCount: result.embeddings.length,
258
+ })
259
+
260
+ if (result.usage) {
261
+ aiEventClient.emit('embedding:usage', {
262
+ requestId,
263
+ model,
264
+ usage: result.usage,
265
+ timestamp: Date.now(),
266
+ })
267
+ await runGenerationUsage(middleware, mwCtx, result.usage)
268
+ }
269
+ await runGenerationFinish(middleware, mwCtx, {
270
+ duration,
271
+ usage: result.usage,
272
+ })
273
+
274
+ return result
275
+ } catch (error) {
276
+ const duration = Date.now() - startTime
277
+ const err = error as Error
278
+ aiEventClient.emit('embedding:request:error', {
279
+ requestId,
280
+ provider: adapter.name,
281
+ model,
282
+ error: { message: err.message, name: err.name },
283
+ duration,
284
+ modelOptions,
285
+ timestamp: Date.now(),
286
+ })
287
+ await runGenerationError(middleware, mwCtx, {
288
+ error,
289
+ duration,
290
+ })
291
+ logger.errors('embed activity failed', {
292
+ error,
293
+ source: 'embed',
294
+ })
295
+ throw error
296
+ }
297
+ }
298
+
299
+ // ===========================
300
+ // Options Factory
301
+ // ===========================
302
+
303
+ /**
304
+ * Create typed options for the embed() function without executing.
305
+ */
306
+ export function createEmbedOptions<
307
+ TAdapter extends EmbeddingAdapter<string, any, any, any>,
308
+ >(options: EmbedOptions<TAdapter>): EmbedOptions<TAdapter> {
309
+ return options
310
+ }
311
+
312
+ // Re-export adapter types
313
+ export type {
314
+ EmbeddingAdapter,
315
+ EmbeddingAdapterConfig,
316
+ AnyEmbeddingAdapter,
317
+ } from './adapter'
318
+ export { BaseEmbeddingAdapter } from './adapter'
@@ -18,6 +18,21 @@ const ABORT_ERROR_NAMES = new Set([
18
18
  'RequestAbortedError',
19
19
  ])
20
20
 
21
+ /**
22
+ * True when a thrown value is an abort-shaped error (DOM `AbortError`, OpenAI
23
+ * `APIUserAbortError`, OpenRouter `RequestAbortedError`) — i.e. user-initiated
24
+ * cancellation rather than a genuine failure. Matches on the error `name` so
25
+ * callers can discriminate aborts without depending on a signal's state or on
26
+ * provider-specific message strings.
27
+ */
28
+ export function isAbortShapedError(error: unknown): boolean {
29
+ if (error && typeof error === 'object') {
30
+ const name = (error as { name?: unknown }).name
31
+ return typeof name === 'string' && ABORT_ERROR_NAMES.has(name)
32
+ }
33
+ return false
34
+ }
35
+
21
36
  // HTTP status codes carried as numbers (e.g. `error.status = 429`) are a
22
37
  // common variant on SDK error classes; coerce so the resulting `code` field
23
38
  // is stable as a string for downstream consumers.
@@ -29,32 +44,49 @@ function normalizeCode(codeField: unknown): string | undefined {
29
44
  return undefined
30
45
  }
31
46
 
47
+ // SDK error classes disagree on where they carry the HTTP status. Most expose a
48
+ // `code` (OpenAI/Anthropic error bodies), but some report it only as a numeric
49
+ // `status` — Google's `@google/genai` `ApiError` sets `status: number` and no
50
+ // `code` at all. Without this fallback such errors reach downstream consumers
51
+ // with `code: undefined`, so a 401/403/404/429 is indistinguishable from an
52
+ // unknown failure and cannot be classified.
53
+ //
54
+ // Only a *numeric* `status` is used: a string `status` is commonly an HTTP
55
+ // reason phrase ("Forbidden") or a symbolic status ("PERMISSION_DENIED"), not
56
+ // the numeric code consumers key on, so forwarding it would be misleading.
57
+ function extractCode(source: {
58
+ code?: unknown
59
+ status?: unknown
60
+ }): string | undefined {
61
+ const fromCode = normalizeCode(source.code)
62
+ if (fromCode !== undefined) return fromCode
63
+ if (typeof source.status === 'number' && Number.isFinite(source.status)) {
64
+ return String(source.status)
65
+ }
66
+ return undefined
67
+ }
68
+
32
69
  export function toRunErrorPayload(
33
70
  error: unknown,
34
71
  fallbackMessage = 'Unknown error occurred',
35
72
  ): { message: string; code: string | undefined } {
36
- if (error && typeof error === 'object') {
37
- const name = (error as { name?: unknown }).name
38
- if (typeof name === 'string' && ABORT_ERROR_NAMES.has(name)) {
39
- return { message: 'Request aborted', code: 'aborted' }
40
- }
73
+ if (isAbortShapedError(error)) {
74
+ return { message: 'Request aborted', code: 'aborted' }
41
75
  }
42
76
  if (error instanceof Error) {
43
- const codeField = (error as Error & { code?: unknown }).code
44
77
  return {
45
78
  message: error.message || fallbackMessage,
46
- code: normalizeCode(codeField),
79
+ code: extractCode(error as Error & { code?: unknown; status?: unknown }),
47
80
  }
48
81
  }
49
82
  if (typeof error === 'object' && error !== null) {
50
83
  const messageField = (error as { message?: unknown }).message
51
- const codeField = (error as { code?: unknown }).code
52
84
  return {
53
85
  message:
54
86
  typeof messageField === 'string' && messageField.length > 0
55
87
  ? messageField
56
88
  : fallbackMessage,
57
- code: normalizeCode(codeField),
89
+ code: extractCode(error as { code?: unknown; status?: unknown }),
58
90
  }
59
91
  }
60
92
  if (typeof error === 'string' && error.length > 0) {
@@ -11,11 +11,18 @@ import { resolveDebugOption } from '../../logger/resolve'
11
11
  import {
12
12
  applyGenerationResultTransforms,
13
13
  createGenerationContext,
14
+ runGenerationAbort,
14
15
  runGenerationError,
15
16
  runGenerationFinish,
16
17
  runGenerationStart,
17
18
  runGenerationUsage,
18
19
  } from '../middleware/run'
20
+ import {
21
+ abortReasonMessage,
22
+ createActivityAbortControls,
23
+ isActivityAbortError,
24
+ raceWithAbort,
25
+ } from '../../utilities/activity-abort'
19
26
  import type { InternalLogger } from '../../logger/internal-logger'
20
27
  import type { DebugOption } from '../../logger/types'
21
28
  import type { GenerationMiddleware } from '../middleware/types'
@@ -89,6 +96,18 @@ export interface AudioActivityOptions<
89
96
  threadId?: string
90
97
  /** Stable run id for correlating this run when persisted. */
91
98
  runId?: string
99
+ /**
100
+ * Maximum duration of this activity invocation in milliseconds.
101
+ * No SDK-wide default — choose a value suitable for the provider and job.
102
+ * Composed with {@link abortSignal}; the first abort wins.
103
+ */
104
+ timeout?: number
105
+ /**
106
+ * Caller cancellation signal (request disconnects, job/runtime cancellation).
107
+ * Composed with {@link timeout} into an effective signal forwarded to the
108
+ * adapter. Request-specific — not stored on global provider client config.
109
+ */
110
+ abortSignal?: AbortSignal
92
111
  }
93
112
 
94
113
  // ===========================
@@ -167,12 +186,18 @@ async function runGenerateAudio<
167
186
  middleware,
168
187
  threadId,
169
188
  runId,
189
+ timeout,
190
+ abortSignal: callerAbortSignal,
170
191
  ...rest
171
192
  } = options
172
193
  const model = adapter.model
173
194
  const requestId = createId('audio')
174
195
  const startTime = Date.now()
175
196
  const logger: InternalLogger = resolveDebugOption(options.debug)
197
+ const abortControls = createActivityAbortControls({
198
+ timeout,
199
+ abortSignal: callerAbortSignal,
200
+ })
176
201
  const providerName =
177
202
  (adapter as { name?: string; provider?: string }).provider ??
178
203
  (adapter as { name?: string }).name ??
@@ -208,7 +233,16 @@ async function runGenerateAudio<
208
233
  })
209
234
 
210
235
  try {
211
- const rawResult = await adapter.generateAudio({ ...rest, model, logger })
236
+ const rawResult = await raceWithAbort(
237
+ adapter.generateAudio({
238
+ ...rest,
239
+ model,
240
+ logger,
241
+ ...(abortControls.signal ? { abortSignal: abortControls.signal } : {}),
242
+ }),
243
+ abortControls.signal,
244
+ )
245
+ abortControls.clear()
212
246
  const result = await applyGenerationResultTransforms(mwCtx, rawResult)
213
247
  const elapsedMs = Date.now() - startTime
214
248
 
@@ -245,6 +279,7 @@ async function runGenerateAudio<
245
279
 
246
280
  return result
247
281
  } catch (error) {
282
+ abortControls.clear()
248
283
  const elapsedMs = Date.now() - startTime
249
284
  const err = error as Error
250
285
  aiEventClient.emit('audio:request:error', {
@@ -256,10 +291,17 @@ async function runGenerateAudio<
256
291
  modelOptions: rest.modelOptions as Record<string, unknown> | undefined,
257
292
  timestamp: Date.now(),
258
293
  })
259
- await runGenerationError(middleware, mwCtx, {
260
- error,
261
- duration: elapsedMs,
262
- })
294
+ if (isActivityAbortError(error, abortControls.signal)) {
295
+ await runGenerationAbort(middleware, mwCtx, {
296
+ reason: abortReasonMessage(error, abortControls.signal),
297
+ duration: elapsedMs,
298
+ })
299
+ } else {
300
+ await runGenerationError(middleware, mwCtx, {
301
+ error,
302
+ duration: elapsedMs,
303
+ })
304
+ }
263
305
  logger.errors('generateAudio activity failed', {
264
306
  error,
265
307
  source: 'generateAudio',
@@ -11,11 +11,18 @@ import { resolveDebugOption } from '../../logger/resolve'
11
11
  import {
12
12
  applyGenerationResultTransforms,
13
13
  createGenerationContext,
14
+ runGenerationAbort,
14
15
  runGenerationError,
15
16
  runGenerationFinish,
16
17
  runGenerationStart,
17
18
  runGenerationUsage,
18
19
  } from '../middleware/run'
20
+ import {
21
+ abortReasonMessage,
22
+ createActivityAbortControls,
23
+ isActivityAbortError,
24
+ raceWithAbort,
25
+ } from '../../utilities/activity-abort'
19
26
  import { resolveMediaPrompt } from '../../utilities/media-prompt'
20
27
  import type { InternalLogger } from '../../logger/internal-logger'
21
28
  import type { DebugOption } from '../../logger/types'
@@ -142,6 +149,18 @@ export type ImageActivityOptions<
142
149
  threadId?: string
143
150
  /** Stable run id for correlating this run when persisted. */
144
151
  runId?: string
152
+ /**
153
+ * Maximum duration of this activity invocation in milliseconds.
154
+ * No SDK-wide default — choose a value suitable for the provider and job.
155
+ * Composed with {@link abortSignal}; the first abort wins.
156
+ */
157
+ timeout?: number
158
+ /**
159
+ * Caller cancellation signal (request disconnects, job/runtime cancellation).
160
+ * Composed with {@link timeout} into an effective signal forwarded to the
161
+ * adapter. Request-specific — not stored on global provider client config.
162
+ */
163
+ abortSignal?: AbortSignal
145
164
  } & ({} extends ImageProviderOptionsForModel<TAdapter, TAdapter['model']>
146
165
  ? {
147
166
  /** Provider-specific options for image generation */ modelOptions?: ImageProviderOptionsForModel<
@@ -260,12 +279,18 @@ async function runGenerateImage<
260
279
  middleware,
261
280
  threadId,
262
281
  runId,
282
+ timeout,
283
+ abortSignal: callerAbortSignal,
263
284
  ...rest
264
285
  } = options
265
286
  const model = adapter.model
266
287
  const requestId = createId('image')
267
288
  const startTime = Date.now()
268
289
  const logger: InternalLogger = resolveDebugOption(options.debug)
290
+ const abortControls = createActivityAbortControls({
291
+ timeout,
292
+ abortSignal: callerAbortSignal,
293
+ })
269
294
 
270
295
  const mwCtx = createGenerationContext({
271
296
  requestId,
@@ -311,7 +336,16 @@ async function runGenerateImage<
311
336
  })
312
337
 
313
338
  try {
314
- const rawResult = await adapter.generateImages({ ...rest, model, logger })
339
+ const rawResult = await raceWithAbort(
340
+ adapter.generateImages({
341
+ ...rest,
342
+ model,
343
+ logger,
344
+ ...(abortControls.signal ? { abortSignal: abortControls.signal } : {}),
345
+ }),
346
+ abortControls.signal,
347
+ )
348
+ abortControls.clear()
315
349
  const result = await applyGenerationResultTransforms(mwCtx, rawResult)
316
350
  const duration = Date.now() - startTime
317
351
 
@@ -355,10 +389,19 @@ async function runGenerateImage<
355
389
 
356
390
  return result
357
391
  } catch (error) {
358
- await runGenerationError(middleware, mwCtx, {
359
- error,
360
- duration: Date.now() - startTime,
361
- })
392
+ abortControls.clear()
393
+ const duration = Date.now() - startTime
394
+ if (isActivityAbortError(error, abortControls.signal)) {
395
+ await runGenerationAbort(middleware, mwCtx, {
396
+ reason: abortReasonMessage(error, abortControls.signal),
397
+ duration,
398
+ })
399
+ } else {
400
+ await runGenerationError(middleware, mwCtx, {
401
+ error,
402
+ duration,
403
+ })
404
+ }
362
405
  logger.errors('generateImage activity failed', {
363
406
  error,
364
407
  source: 'generateImage',