ai 7.0.92 → 7.0.94
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/CHANGELOG.md +33 -0
- package/dist/index.d.ts +239 -214
- package/dist/index.js +391 -138
- package/dist/index.js.map +1 -1
- package/dist/internal/index.d.ts +3 -2
- package/dist/internal/index.js +8 -5
- package/dist/internal/index.js.map +1 -1
- package/docs/02-foundations/02-providers-and-models.mdx +1 -1
- package/docs/03-agents/04-loop-control.mdx +5 -3
- package/docs/03-agents/07-workflow-agent.mdx +27 -6
- package/docs/03-ai-sdk-core/10-generating-structured-data.mdx +15 -3
- package/docs/03-ai-sdk-core/16-mcp-tools.mdx +64 -1
- package/docs/03-ai-sdk-core/35-image-generation.mdx +9 -0
- package/docs/03-ai-sdk-core/36-transcription.mdx +36 -35
- package/docs/03-ai-sdk-core/37-speech.mdx +0 -17
- package/docs/03-ai-sdk-harnesses/02-harness-agent.mdx +36 -0
- package/docs/04-ai-sdk-ui/20-streaming-data.mdx +11 -6
- package/docs/07-reference/01-ai-sdk-core/01-generate-text.mdx +14 -0
- package/docs/07-reference/01-ai-sdk-core/02-stream-text.mdx +15 -1
- package/docs/07-reference/01-ai-sdk-core/10-generate-image.mdx +2 -1
- package/docs/07-reference/01-ai-sdk-core/28-output.mdx +27 -1
- package/docs/07-reference/02-ai-sdk-ui/01-use-chat.mdx +1 -1
- package/docs/07-reference/02-ai-sdk-ui/40-create-ui-message-stream.mdx +4 -0
- package/docs/07-reference/02-ai-sdk-ui/41-create-ui-message-stream-response.mdx +6 -1
- package/docs/07-reference/04-ai-sdk-workflow/01-workflow-agent.mdx +42 -28
- package/docs/07-reference/05-ai-sdk-errors/ai-no-image-generated-error.mdx +7 -0
- package/package.json +12 -12
- package/src/agent/tool-loop-agent-settings.ts +15 -0
- package/src/batch/batch-types.ts +76 -68
- package/src/batch/batch.ts +169 -101
- package/src/batch/index.ts +9 -7
- package/src/embed/embed-many.ts +27 -2
- package/src/error/no-image-generated-error.ts +9 -0
- package/src/generate-image/generate-image.ts +51 -10
- package/src/generate-text/output.ts +111 -1
- package/src/generate-text/stream-language-model-call.ts +82 -0
- package/src/generate-text/stream-text.ts +5 -1
- package/src/ui/chat.ts +1 -1
- package/src/ui/convert-to-model-messages.ts +8 -2
- package/src/ui/validate-ui-messages.ts +14 -0
- package/src/util/data-url.ts +1 -1
- package/src/util/merge-abort-signals.ts +1 -1
- package/src/util/prepare-retries.ts +7 -1
- package/src/util/retry-with-exponential-backoff.ts +9 -4
package/src/batch/batch.ts
CHANGED
|
@@ -1,11 +1,12 @@
|
|
|
1
1
|
import {
|
|
2
2
|
UnsupportedFunctionalityError,
|
|
3
|
-
type
|
|
3
|
+
type Experimental_BatchV4 as BatchV4,
|
|
4
4
|
type Experimental_BatchV4ItemResult as BatchV4ItemResult,
|
|
5
|
-
type LanguageModelV4,
|
|
6
5
|
type LanguageModelV4GenerateResult,
|
|
7
6
|
type LanguageModelV4ToolCall,
|
|
7
|
+
type ProviderV4,
|
|
8
8
|
} from '@ai-sdk/provider';
|
|
9
|
+
import { gateway } from '@ai-sdk/gateway';
|
|
9
10
|
import { type ToolSet, withUserAgentSuffix } from '@ai-sdk/provider-utils';
|
|
10
11
|
import { InvalidArgumentError } from '../error/invalid-argument-error';
|
|
11
12
|
import { convertLanguageModelContent } from '../generate-text/convert-language-model-content';
|
|
@@ -13,7 +14,6 @@ import { parseToolCall } from '../generate-text/parse-tool-call';
|
|
|
13
14
|
import { prepareToolChoice } from '../prompt/prepare-tool-choice';
|
|
14
15
|
import { prepareTools } from '../prompt/prepare-tools';
|
|
15
16
|
import { logWarnings } from '../logger/log-warnings';
|
|
16
|
-
import { resolveLanguageModel } from '../model/resolve-model';
|
|
17
17
|
import { convertToLanguageModelPrompt } from '../prompt/convert-to-language-model-prompt';
|
|
18
18
|
import { prepareLanguageModelCallOptions } from '../prompt/prepare-language-model-call-options';
|
|
19
19
|
import { getTotalTimeoutMs } from '../prompt/request-options';
|
|
@@ -22,70 +22,85 @@ import { wrapGatewayError } from '../prompt/wrap-gateway-error';
|
|
|
22
22
|
import { asLanguageModelUsage } from '../types/usage';
|
|
23
23
|
import { asAsyncIterableStream } from '../util/async-iterable-stream';
|
|
24
24
|
import { mergeAbortSignals } from '../util/merge-abort-signals';
|
|
25
|
+
import { isDeepEqualData } from '../util/is-deep-equal-data';
|
|
25
26
|
import { prepareRetries } from '../util/prepare-retries';
|
|
26
27
|
import { VERSION } from '../version';
|
|
28
|
+
import { asProviderV4 } from '../model/as-provider-v4';
|
|
27
29
|
import type {
|
|
28
|
-
|
|
30
|
+
BatchItemResult,
|
|
31
|
+
BatchProvider,
|
|
29
32
|
BatchReference,
|
|
30
33
|
BatchStatus,
|
|
31
|
-
|
|
32
|
-
|
|
34
|
+
GetBatchResultsOptions,
|
|
35
|
+
GetBatchStatusOptions,
|
|
36
|
+
StartBatchOptions,
|
|
37
|
+
StartBatchResult,
|
|
33
38
|
TextBatchGenerationResult,
|
|
34
39
|
TextBatchItemResult,
|
|
35
|
-
TextBatchRequest,
|
|
36
40
|
} from './batch-types';
|
|
37
41
|
|
|
38
42
|
/**
|
|
39
|
-
* Starts a
|
|
43
|
+
* Starts a batch.
|
|
40
44
|
*/
|
|
41
|
-
export async function
|
|
42
|
-
|
|
45
|
+
export async function startBatch<
|
|
46
|
+
TOOLS extends ToolSet,
|
|
47
|
+
PROVIDER extends BatchProvider,
|
|
48
|
+
>({
|
|
49
|
+
provider,
|
|
43
50
|
requests,
|
|
44
|
-
tools,
|
|
45
|
-
toolChoice,
|
|
46
|
-
toolOrder,
|
|
47
|
-
toolsContext,
|
|
48
51
|
providerOptions,
|
|
49
52
|
webhookUrl,
|
|
50
53
|
abortSignal,
|
|
51
54
|
headers,
|
|
52
55
|
timeout,
|
|
53
|
-
}:
|
|
56
|
+
}: StartBatchOptions<TOOLS, PROVIDER>): Promise<StartBatchResult> {
|
|
54
57
|
validateRequests(requests);
|
|
55
58
|
|
|
56
|
-
const
|
|
59
|
+
const batchApi = resolveBatchApi(provider);
|
|
57
60
|
const operationAbortSignal = mergeAbortSignals(
|
|
58
61
|
abortSignal,
|
|
59
62
|
getTotalTimeoutMs(timeout),
|
|
60
63
|
);
|
|
61
|
-
const supportedUrls = await
|
|
62
|
-
const preparedTools = await prepareTools({
|
|
63
|
-
tools,
|
|
64
|
-
toolOrder,
|
|
65
|
-
toolsContext,
|
|
66
|
-
});
|
|
67
|
-
const preparedToolChoice = prepareToolChoice({ toolChoice });
|
|
64
|
+
const supportedUrls = await batchApi.supportedUrls;
|
|
68
65
|
operationAbortSignal?.throwIfAborted();
|
|
69
66
|
const normalizedRequests = [];
|
|
67
|
+
const toolsByName = new Map<string, unknown>();
|
|
70
68
|
|
|
71
69
|
for (const request of requests) {
|
|
72
|
-
|
|
73
|
-
|
|
74
|
-
|
|
75
|
-
|
|
76
|
-
|
|
77
|
-
|
|
78
|
-
|
|
79
|
-
|
|
80
|
-
|
|
81
|
-
|
|
82
|
-
|
|
83
|
-
|
|
84
|
-
|
|
85
|
-
|
|
86
|
-
|
|
87
|
-
|
|
88
|
-
|
|
70
|
+
switch (request.type) {
|
|
71
|
+
case 'text': {
|
|
72
|
+
const standardizedPrompt = await standardizePrompt(request);
|
|
73
|
+
const preparedTools = await prepareTools({
|
|
74
|
+
tools: request.tools,
|
|
75
|
+
toolOrder: request.toolOrder,
|
|
76
|
+
toolsContext: request.toolsContext,
|
|
77
|
+
});
|
|
78
|
+
validateCompatibleTools({
|
|
79
|
+
requestId: request.id,
|
|
80
|
+
tools: preparedTools,
|
|
81
|
+
toolsByName,
|
|
82
|
+
});
|
|
83
|
+
|
|
84
|
+
normalizedRequests.push({
|
|
85
|
+
id: request.id,
|
|
86
|
+
type: request.type,
|
|
87
|
+
modelId: request.model,
|
|
88
|
+
options: {
|
|
89
|
+
...prepareLanguageModelCallOptions(request),
|
|
90
|
+
prompt: await convertToLanguageModelPrompt({
|
|
91
|
+
prompt: standardizedPrompt,
|
|
92
|
+
supportedUrls,
|
|
93
|
+
download: undefined,
|
|
94
|
+
provider: batchApi.provider.split('.')[0],
|
|
95
|
+
}),
|
|
96
|
+
tools: preparedTools,
|
|
97
|
+
toolChoice: prepareToolChoice({ toolChoice: request.toolChoice }),
|
|
98
|
+
providerOptions: request.providerOptions,
|
|
99
|
+
},
|
|
100
|
+
});
|
|
101
|
+
break;
|
|
102
|
+
}
|
|
103
|
+
}
|
|
89
104
|
operationAbortSignal?.throwIfAborted();
|
|
90
105
|
}
|
|
91
106
|
|
|
@@ -94,7 +109,7 @@ export async function startTextBatch<TOOLS extends ToolSet>({
|
|
|
94
109
|
`ai/${VERSION}`,
|
|
95
110
|
);
|
|
96
111
|
try {
|
|
97
|
-
const result = await
|
|
112
|
+
const result = await batchApi.doStartBatch({
|
|
98
113
|
requests: normalizedRequests,
|
|
99
114
|
providerOptions,
|
|
100
115
|
abortSignal: operationAbortSignal,
|
|
@@ -103,18 +118,21 @@ export async function startTextBatch<TOOLS extends ToolSet>({
|
|
|
103
118
|
});
|
|
104
119
|
const { batchId, warnings, ...status } = result;
|
|
105
120
|
|
|
106
|
-
|
|
107
|
-
|
|
108
|
-
|
|
109
|
-
|
|
110
|
-
|
|
121
|
+
const modelByRequestId = new Map(
|
|
122
|
+
normalizedRequests.map(request => [request.id, request.modelId]),
|
|
123
|
+
);
|
|
124
|
+
for (const { requestId, warning } of warnings) {
|
|
125
|
+
logWarnings({
|
|
126
|
+
warnings: [warning],
|
|
127
|
+
provider: batchApi.provider,
|
|
128
|
+
model: requestId == null ? undefined : modelByRequestId.get(requestId),
|
|
129
|
+
});
|
|
130
|
+
}
|
|
111
131
|
|
|
112
132
|
return {
|
|
113
|
-
version:
|
|
114
|
-
type: 'text',
|
|
133
|
+
version: 2,
|
|
115
134
|
id: batchId,
|
|
116
|
-
provider:
|
|
117
|
-
modelId: model.modelId,
|
|
135
|
+
provider: batchApi.provider,
|
|
118
136
|
...status,
|
|
119
137
|
warnings,
|
|
120
138
|
};
|
|
@@ -123,20 +141,44 @@ export async function startTextBatch<TOOLS extends ToolSet>({
|
|
|
123
141
|
}
|
|
124
142
|
}
|
|
125
143
|
|
|
144
|
+
function validateCompatibleTools({
|
|
145
|
+
requestId,
|
|
146
|
+
tools,
|
|
147
|
+
toolsByName,
|
|
148
|
+
}: {
|
|
149
|
+
requestId: string;
|
|
150
|
+
tools: ReadonlyArray<{ name: string }> | undefined;
|
|
151
|
+
toolsByName: Map<string, unknown>;
|
|
152
|
+
}) {
|
|
153
|
+
for (const tool of tools ?? []) {
|
|
154
|
+
const previousTool = toolsByName.get(tool.name);
|
|
155
|
+
|
|
156
|
+
if (previousTool != null && !isDeepEqualData(previousTool, tool)) {
|
|
157
|
+
throw new InvalidArgumentError({
|
|
158
|
+
parameter: 'requests',
|
|
159
|
+
value: requestId,
|
|
160
|
+
message: `tool "${tool.name}" must have the same definition in every batch request`,
|
|
161
|
+
});
|
|
162
|
+
}
|
|
163
|
+
|
|
164
|
+
toolsByName.set(tool.name, tool);
|
|
165
|
+
}
|
|
166
|
+
}
|
|
167
|
+
|
|
126
168
|
/**
|
|
127
|
-
* Retrieves the latest normalized status for a
|
|
169
|
+
* Retrieves the latest normalized status for a batch.
|
|
128
170
|
*/
|
|
129
171
|
export async function getBatchStatus({
|
|
130
|
-
|
|
172
|
+
provider,
|
|
131
173
|
batch,
|
|
132
174
|
providerOptions,
|
|
133
175
|
maxRetries,
|
|
134
176
|
abortSignal,
|
|
135
177
|
headers,
|
|
136
178
|
timeout,
|
|
137
|
-
}:
|
|
138
|
-
const
|
|
139
|
-
validateBatchReference({
|
|
179
|
+
}: GetBatchStatusOptions): Promise<BatchStatus> {
|
|
180
|
+
const batchApi = resolveBatchApi(provider);
|
|
181
|
+
validateBatchReference({ batchApi, batch });
|
|
140
182
|
|
|
141
183
|
const operationAbortSignal = mergeAbortSignals(
|
|
142
184
|
abortSignal,
|
|
@@ -149,7 +191,7 @@ export async function getBatchStatus({
|
|
|
149
191
|
|
|
150
192
|
try {
|
|
151
193
|
const status = await retry(() =>
|
|
152
|
-
|
|
194
|
+
batchApi.doGetBatchStatus({
|
|
153
195
|
batchId: batch.id,
|
|
154
196
|
providerOptions,
|
|
155
197
|
abortSignal: operationAbortSignal,
|
|
@@ -164,10 +206,10 @@ export async function getBatchStatus({
|
|
|
164
206
|
}
|
|
165
207
|
|
|
166
208
|
/**
|
|
167
|
-
* Streams complete terminal results for the requests in a
|
|
209
|
+
* Streams complete terminal results for the requests in a batch.
|
|
168
210
|
*/
|
|
169
211
|
export function getBatchResults<TOOLS extends ToolSet>({
|
|
170
|
-
|
|
212
|
+
provider,
|
|
171
213
|
batch,
|
|
172
214
|
tools,
|
|
173
215
|
providerOptions,
|
|
@@ -175,9 +217,9 @@ export function getBatchResults<TOOLS extends ToolSet>({
|
|
|
175
217
|
abortSignal,
|
|
176
218
|
headers,
|
|
177
219
|
timeout,
|
|
178
|
-
}:
|
|
179
|
-
const
|
|
180
|
-
validateBatchReference({
|
|
220
|
+
}: GetBatchResultsOptions<TOOLS>) {
|
|
221
|
+
const batchApi = resolveBatchApi(provider);
|
|
222
|
+
validateBatchReference({ batchApi, batch });
|
|
181
223
|
|
|
182
224
|
const streamAbortController = new AbortController();
|
|
183
225
|
const operationAbortSignal = mergeAbortSignals(
|
|
@@ -189,10 +231,9 @@ export function getBatchResults<TOOLS extends ToolSet>({
|
|
|
189
231
|
maxRetries,
|
|
190
232
|
abortSignal: operationAbortSignal,
|
|
191
233
|
});
|
|
192
|
-
const transformer: Transformer<
|
|
193
|
-
|
|
194
|
-
|
|
195
|
-
> & { cancel?: (reason?: unknown) => void } = {
|
|
234
|
+
const transformer: Transformer<BatchV4ItemResult, BatchItemResult<TOOLS>> & {
|
|
235
|
+
cancel?: (reason?: unknown) => void;
|
|
236
|
+
} = {
|
|
196
237
|
async transform(item, controller) {
|
|
197
238
|
controller.enqueue(await convertBatchItemResult({ item, tools }));
|
|
198
239
|
},
|
|
@@ -204,14 +245,14 @@ export function getBatchResults<TOOLS extends ToolSet>({
|
|
|
204
245
|
},
|
|
205
246
|
};
|
|
206
247
|
const transform = new TransformStream<
|
|
207
|
-
BatchV4ItemResult
|
|
208
|
-
|
|
248
|
+
BatchV4ItemResult,
|
|
249
|
+
BatchItemResult<TOOLS>
|
|
209
250
|
>(transformer);
|
|
210
251
|
|
|
211
252
|
void (async () => {
|
|
212
253
|
try {
|
|
213
254
|
const stream = await retry(() =>
|
|
214
|
-
|
|
255
|
+
batchApi.doGetBatchResults({
|
|
215
256
|
batchId: batch.id,
|
|
216
257
|
providerOptions,
|
|
217
258
|
abortSignal: operationAbortSignal,
|
|
@@ -230,33 +271,43 @@ export function getBatchResults<TOOLS extends ToolSet>({
|
|
|
230
271
|
return asAsyncIterableStream(transform.readable);
|
|
231
272
|
}
|
|
232
273
|
|
|
233
|
-
function
|
|
234
|
-
|
|
235
|
-
|
|
236
|
-
|
|
274
|
+
function resolveBatchApi(provider?: BatchProvider): BatchV4 {
|
|
275
|
+
provider ??= asProviderV4(globalThis.AI_SDK_DEFAULT_PROVIDER ?? gateway);
|
|
276
|
+
|
|
277
|
+
if (isBatchApi(provider)) {
|
|
278
|
+
return provider;
|
|
279
|
+
}
|
|
237
280
|
|
|
238
|
-
if (!
|
|
281
|
+
if (!hasBatchFactory(provider)) {
|
|
239
282
|
throw new UnsupportedFunctionalityError({
|
|
240
283
|
functionality: 'batch processing',
|
|
241
|
-
message:
|
|
284
|
+
message:
|
|
285
|
+
'The provider does not support batch processing. Make sure it exposes an experimental_batch() method.',
|
|
242
286
|
});
|
|
243
287
|
}
|
|
244
288
|
|
|
245
|
-
return
|
|
289
|
+
return provider.experimental_batch();
|
|
290
|
+
}
|
|
291
|
+
|
|
292
|
+
function hasBatchFactory(
|
|
293
|
+
provider: ProviderV4,
|
|
294
|
+
): provider is ProviderV4 & { experimental_batch(): BatchV4 } {
|
|
295
|
+
return (
|
|
296
|
+
typeof (provider as { experimental_batch?: unknown }).experimental_batch ===
|
|
297
|
+
'function'
|
|
298
|
+
);
|
|
246
299
|
}
|
|
247
300
|
|
|
248
|
-
function
|
|
249
|
-
|
|
250
|
-
): model is BatchLanguageModelV4 {
|
|
251
|
-
const candidate = model as Partial<BatchLanguageModelV4>;
|
|
301
|
+
function isBatchApi(provider: BatchProvider): provider is BatchV4 {
|
|
302
|
+
const candidate = provider as Partial<BatchV4>;
|
|
252
303
|
return (
|
|
253
|
-
typeof candidate.
|
|
254
|
-
typeof candidate.
|
|
255
|
-
typeof candidate.
|
|
304
|
+
typeof candidate.doStartBatch === 'function' &&
|
|
305
|
+
typeof candidate.doGetBatchStatus === 'function' &&
|
|
306
|
+
typeof candidate.doGetBatchResults === 'function'
|
|
256
307
|
);
|
|
257
308
|
}
|
|
258
309
|
|
|
259
|
-
function validateRequests(requests: ReadonlyArray<
|
|
310
|
+
function validateRequests(requests: ReadonlyArray<{ readonly id: string }>) {
|
|
260
311
|
if (requests.length === 0) {
|
|
261
312
|
throw new InvalidArgumentError({
|
|
262
313
|
parameter: 'requests',
|
|
@@ -289,27 +340,27 @@ function validateRequests(requests: ReadonlyArray<TextBatchRequest>) {
|
|
|
289
340
|
}
|
|
290
341
|
|
|
291
342
|
function validateBatchReference({
|
|
292
|
-
|
|
343
|
+
batchApi,
|
|
293
344
|
batch,
|
|
294
345
|
}: {
|
|
295
|
-
|
|
346
|
+
batchApi: BatchV4;
|
|
296
347
|
batch: BatchReference;
|
|
297
348
|
}) {
|
|
298
|
-
if (batch.version !==
|
|
349
|
+
if (batch.version !== 2) {
|
|
299
350
|
throw new InvalidArgumentError({
|
|
300
351
|
parameter: 'batch',
|
|
301
352
|
value: batch,
|
|
302
|
-
message: 'batch must be a supported
|
|
353
|
+
message: 'batch must be a supported batch reference',
|
|
303
354
|
});
|
|
304
355
|
}
|
|
305
356
|
|
|
306
|
-
if (batch.provider !==
|
|
357
|
+
if (batch.provider !== batchApi.provider) {
|
|
307
358
|
throw new InvalidArgumentError({
|
|
308
|
-
parameter: '
|
|
309
|
-
value:
|
|
359
|
+
parameter: 'provider',
|
|
360
|
+
value: batchApi,
|
|
310
361
|
message:
|
|
311
|
-
`
|
|
312
|
-
`batch ${batch.provider}
|
|
362
|
+
`provider ${batchApi.provider} is not compatible with ` +
|
|
363
|
+
`batch provider ${batch.provider}`,
|
|
313
364
|
});
|
|
314
365
|
}
|
|
315
366
|
}
|
|
@@ -318,18 +369,35 @@ async function convertBatchItemResult<TOOLS extends ToolSet>({
|
|
|
318
369
|
item,
|
|
319
370
|
tools,
|
|
320
371
|
}: {
|
|
321
|
-
item: BatchV4ItemResult
|
|
372
|
+
item: BatchV4ItemResult;
|
|
322
373
|
tools: TOOLS | undefined;
|
|
323
374
|
}): Promise<TextBatchItemResult<TOOLS>> {
|
|
324
|
-
|
|
325
|
-
|
|
375
|
+
switch (item.type) {
|
|
376
|
+
case 'text':
|
|
377
|
+
switch (item.status) {
|
|
378
|
+
case 'succeeded':
|
|
379
|
+
return {
|
|
380
|
+
id: item.id,
|
|
381
|
+
status: item.status,
|
|
382
|
+
...(await convertGenerateResult({ result: item.result, tools })),
|
|
383
|
+
};
|
|
384
|
+
case 'failed':
|
|
385
|
+
return {
|
|
386
|
+
id: item.id,
|
|
387
|
+
status: item.status,
|
|
388
|
+
error: item.error,
|
|
389
|
+
providerMetadata: item.providerMetadata,
|
|
390
|
+
};
|
|
391
|
+
case 'cancelled':
|
|
392
|
+
case 'expired':
|
|
393
|
+
return {
|
|
394
|
+
id: item.id,
|
|
395
|
+
status: item.status,
|
|
396
|
+
error: item.error,
|
|
397
|
+
providerMetadata: item.providerMetadata,
|
|
398
|
+
};
|
|
399
|
+
}
|
|
326
400
|
}
|
|
327
|
-
|
|
328
|
-
return {
|
|
329
|
-
id: item.id,
|
|
330
|
-
status: 'succeeded',
|
|
331
|
-
...(await convertGenerateResult({ result: item.result, tools })),
|
|
332
|
-
};
|
|
333
401
|
}
|
|
334
402
|
|
|
335
403
|
async function convertGenerateResult<TOOLS extends ToolSet>({
|
package/src/batch/index.ts
CHANGED
|
@@ -1,19 +1,21 @@
|
|
|
1
1
|
export {
|
|
2
|
-
|
|
2
|
+
startBatch as experimental_startBatch,
|
|
3
3
|
getBatchResults as experimental_getBatchResults,
|
|
4
4
|
getBatchStatus as experimental_getBatchStatus,
|
|
5
5
|
} from './batch';
|
|
6
6
|
export type {
|
|
7
7
|
BatchError as Experimental_BatchError,
|
|
8
|
-
|
|
9
|
-
BatchOperationOptions as Experimental_BatchOperationOptions,
|
|
8
|
+
BatchProvider as Experimental_BatchProvider,
|
|
10
9
|
BatchReference as Experimental_BatchReference,
|
|
11
10
|
BatchStatus as Experimental_BatchStatus,
|
|
12
|
-
|
|
13
|
-
|
|
14
|
-
|
|
11
|
+
GetBatchResultsOptions as Experimental_GetBatchResultsOptions,
|
|
12
|
+
GetBatchStatusOptions as Experimental_GetBatchStatusOptions,
|
|
13
|
+
StartBatchOptions as Experimental_StartBatchOptions,
|
|
14
|
+
StartBatchResult as Experimental_StartBatchResult,
|
|
15
|
+
Batch as Experimental_Batch,
|
|
16
|
+
BatchItemResult as Experimental_BatchItemResult,
|
|
17
|
+
BatchRequest as Experimental_BatchRequest,
|
|
15
18
|
TextBatchGenerationResult as Experimental_TextBatchGenerationResult,
|
|
16
19
|
TextBatchItemResult as Experimental_TextBatchItemResult,
|
|
17
|
-
TextBatchReference as Experimental_TextBatchReference,
|
|
18
20
|
TextBatchRequest as Experimental_TextBatchRequest,
|
|
19
21
|
} from './batch-types';
|
package/src/embed/embed-many.ts
CHANGED
|
@@ -1,3 +1,4 @@
|
|
|
1
|
+
import { InvalidResponseDataError } from '@ai-sdk/provider';
|
|
1
2
|
import {
|
|
2
3
|
createIdGenerator,
|
|
3
4
|
withUserAgentSuffix,
|
|
@@ -264,6 +265,8 @@ export async function embedMany({
|
|
|
264
265
|
};
|
|
265
266
|
});
|
|
266
267
|
|
|
268
|
+
validateEmbeddingCount({ embeddings, values });
|
|
269
|
+
|
|
267
270
|
logWarnings({
|
|
268
271
|
warnings,
|
|
269
272
|
provider: model.provider,
|
|
@@ -325,8 +328,8 @@ export async function embedMany({
|
|
|
325
328
|
|
|
326
329
|
for (const parallelChunk of parallelChunks) {
|
|
327
330
|
const results = await Promise.all(
|
|
328
|
-
parallelChunk.map(chunk => {
|
|
329
|
-
|
|
331
|
+
parallelChunk.map(async chunk => {
|
|
332
|
+
const result = await retry(async () => {
|
|
330
333
|
const embedCallId = generateCallId();
|
|
331
334
|
|
|
332
335
|
await notify({
|
|
@@ -373,6 +376,13 @@ export async function embedMany({
|
|
|
373
376
|
response: modelResponse.response,
|
|
374
377
|
};
|
|
375
378
|
});
|
|
379
|
+
|
|
380
|
+
validateEmbeddingCount({
|
|
381
|
+
embeddings: result.embeddings,
|
|
382
|
+
values: chunk,
|
|
383
|
+
});
|
|
384
|
+
|
|
385
|
+
return result;
|
|
376
386
|
}),
|
|
377
387
|
);
|
|
378
388
|
|
|
@@ -436,6 +446,21 @@ export async function embedMany({
|
|
|
436
446
|
});
|
|
437
447
|
}
|
|
438
448
|
|
|
449
|
+
function validateEmbeddingCount({
|
|
450
|
+
embeddings,
|
|
451
|
+
values,
|
|
452
|
+
}: {
|
|
453
|
+
embeddings: Array<Embedding>;
|
|
454
|
+
values: Array<string>;
|
|
455
|
+
}) {
|
|
456
|
+
if (embeddings.length !== values.length) {
|
|
457
|
+
throw new InvalidResponseDataError({
|
|
458
|
+
data: embeddings,
|
|
459
|
+
message: `Expected ${values.length} embeddings, but received ${embeddings.length}.`,
|
|
460
|
+
});
|
|
461
|
+
}
|
|
462
|
+
}
|
|
463
|
+
|
|
439
464
|
const textEncoder = new TextEncoder();
|
|
440
465
|
|
|
441
466
|
function splitByEmbeddingLimits({
|
|
@@ -1,4 +1,5 @@
|
|
|
1
1
|
import { AISDKError } from '@ai-sdk/provider';
|
|
2
|
+
import type { GenerateImageCall } from '../generate-image/generate-image-result';
|
|
2
3
|
import type { ImageModelResponseMetadata } from '../types/image-model-response-metadata';
|
|
3
4
|
|
|
4
5
|
const name = 'AI_NoImageGeneratedError';
|
|
@@ -14,6 +15,11 @@ const symbol = Symbol.for(marker);
|
|
|
14
15
|
export class NoImageGeneratedError extends AISDKError {
|
|
15
16
|
private readonly [symbol] = true; // used in isInstance
|
|
16
17
|
|
|
18
|
+
/**
|
|
19
|
+
* The results of the underlying image model calls.
|
|
20
|
+
*/
|
|
21
|
+
readonly calls: Array<GenerateImageCall> | undefined;
|
|
22
|
+
|
|
17
23
|
/**
|
|
18
24
|
* The response metadata for each call.
|
|
19
25
|
*/
|
|
@@ -22,14 +28,17 @@ export class NoImageGeneratedError extends AISDKError {
|
|
|
22
28
|
constructor({
|
|
23
29
|
message = 'No image generated.',
|
|
24
30
|
cause,
|
|
31
|
+
calls,
|
|
25
32
|
responses,
|
|
26
33
|
}: {
|
|
27
34
|
message?: string;
|
|
28
35
|
cause?: Error;
|
|
36
|
+
calls?: Array<GenerateImageCall>;
|
|
29
37
|
responses?: Array<ImageModelResponseMetadata>;
|
|
30
38
|
}) {
|
|
31
39
|
super({ name, message, cause });
|
|
32
40
|
|
|
41
|
+
this.calls = calls;
|
|
33
42
|
this.responses = responses;
|
|
34
43
|
}
|
|
35
44
|
|
|
@@ -4,6 +4,7 @@ import {
|
|
|
4
4
|
type ImageModelV4CallOptions,
|
|
5
5
|
type ImageModelV4File,
|
|
6
6
|
type ImageModelV4ProviderMetadata,
|
|
7
|
+
type ImageModelV4Result,
|
|
7
8
|
type JSONObject,
|
|
8
9
|
} from '@ai-sdk/provider';
|
|
9
10
|
import {
|
|
@@ -25,6 +26,7 @@ import type { ImageModelResponseMetadata } from '../types/image-model-response-m
|
|
|
25
26
|
import { addImageModelUsage, type ImageModelUsage } from '../types/usage';
|
|
26
27
|
import type { Warning } from '../types/warning';
|
|
27
28
|
import { prepareRetries } from '../util/prepare-retries';
|
|
29
|
+
import { RetryError } from '../util/retry-error';
|
|
28
30
|
import { VERSION } from '../version';
|
|
29
31
|
import type {
|
|
30
32
|
GenerateImageCall,
|
|
@@ -47,6 +49,13 @@ type GatewayCostMetadata = {
|
|
|
47
49
|
[key in (typeof gatewayCostMetadataKeys)[number]]?: unknown;
|
|
48
50
|
};
|
|
49
51
|
|
|
52
|
+
class RetryableNoImageResultError extends Error {
|
|
53
|
+
constructor() {
|
|
54
|
+
super('No image generated.');
|
|
55
|
+
this.name = 'RetryableNoImageResultError';
|
|
56
|
+
}
|
|
57
|
+
}
|
|
58
|
+
|
|
50
59
|
export type GenerateImagePrompt =
|
|
51
60
|
| string
|
|
52
61
|
| {
|
|
@@ -67,7 +76,7 @@ export type GenerateImagePrompt =
|
|
|
67
76
|
* @param seed - Seed for the image generation.
|
|
68
77
|
* @param providerOptions - Additional provider-specific options that are passed through to the provider
|
|
69
78
|
* as body parameters.
|
|
70
|
-
* @param maxRetries - Maximum number of retries. Set to 0 to disable retries. Default: 2.
|
|
79
|
+
* @param maxRetries - Maximum number of retries per image model call, including retries after unclassified empty responses. Empty responses marked as not retryable by the provider are not retried. Set to 0 to disable retries. Default: 2.
|
|
71
80
|
* @param abortSignal - An optional abort signal that can be used to cancel the call.
|
|
72
81
|
* @param headers - Additional HTTP headers to be sent with the request. Only applicable for HTTP-based providers.
|
|
73
82
|
*
|
|
@@ -138,7 +147,9 @@ export async function generateImage({
|
|
|
138
147
|
providerOptions?: ProviderOptions;
|
|
139
148
|
|
|
140
149
|
/**
|
|
141
|
-
* Maximum number of retries per image model call
|
|
150
|
+
* Maximum number of retries per image model call, including retries after
|
|
151
|
+
* unclassified empty responses. Empty responses marked as not retryable by
|
|
152
|
+
* the provider are not retried. Set to 0 to disable retries.
|
|
142
153
|
*
|
|
143
154
|
* @default 2
|
|
144
155
|
*/
|
|
@@ -165,6 +176,8 @@ export async function generateImage({
|
|
|
165
176
|
const { retry } = prepareRetries({
|
|
166
177
|
maxRetries: maxRetriesArg,
|
|
167
178
|
abortSignal,
|
|
179
|
+
additionalRetryableError: error =>
|
|
180
|
+
error instanceof RetryableNoImageResultError,
|
|
168
181
|
});
|
|
169
182
|
|
|
170
183
|
// default to 1 if the model has not specified limits on
|
|
@@ -183,13 +196,15 @@ export async function generateImage({
|
|
|
183
196
|
return remainder === 0 ? maxImagesPerCallWithDefault : remainder;
|
|
184
197
|
});
|
|
185
198
|
|
|
186
|
-
const
|
|
187
|
-
callImageCounts.map(
|
|
188
|
-
|
|
189
|
-
|
|
199
|
+
const resultGroups = await Promise.all(
|
|
200
|
+
callImageCounts.map(async callImageCount => {
|
|
201
|
+
const callResults: Array<ImageModelV4Result> = [];
|
|
202
|
+
|
|
203
|
+
try {
|
|
204
|
+
await retry(async () => {
|
|
190
205
|
const { prompt, files, mask } = normalizePrompt(promptArg);
|
|
191
206
|
|
|
192
|
-
|
|
207
|
+
const result = await model.doGenerate({
|
|
193
208
|
prompt,
|
|
194
209
|
files,
|
|
195
210
|
mask,
|
|
@@ -201,9 +216,35 @@ export async function generateImage({
|
|
|
201
216
|
seed,
|
|
202
217
|
providerOptions: providerOptions ?? {},
|
|
203
218
|
});
|
|
204
|
-
|
|
205
|
-
|
|
219
|
+
|
|
220
|
+
callResults.push(result);
|
|
221
|
+
|
|
222
|
+
if (result.images.length === 0 && result.isRetryable !== false) {
|
|
223
|
+
throw new RetryableNoImageResultError();
|
|
224
|
+
}
|
|
225
|
+
|
|
226
|
+
return result;
|
|
227
|
+
});
|
|
228
|
+
|
|
229
|
+
return callResults;
|
|
230
|
+
} catch (error) {
|
|
231
|
+
const noImageResultError =
|
|
232
|
+
error instanceof RetryableNoImageResultError
|
|
233
|
+
? error
|
|
234
|
+
: RetryError.isInstance(error) &&
|
|
235
|
+
error.lastError instanceof RetryableNoImageResultError
|
|
236
|
+
? error.lastError
|
|
237
|
+
: undefined;
|
|
238
|
+
|
|
239
|
+
if (noImageResultError != null) {
|
|
240
|
+
return callResults;
|
|
241
|
+
}
|
|
242
|
+
|
|
243
|
+
throw error;
|
|
244
|
+
}
|
|
245
|
+
}),
|
|
206
246
|
);
|
|
247
|
+
const results = resultGroups.flat();
|
|
207
248
|
|
|
208
249
|
// collect result images, warnings, and response metadata
|
|
209
250
|
const images: Array<GeneratedFile> = [];
|
|
@@ -295,7 +336,7 @@ export async function generateImage({
|
|
|
295
336
|
logWarnings({ warnings, provider: model.provider, model: model.modelId });
|
|
296
337
|
|
|
297
338
|
if (!images.length) {
|
|
298
|
-
throw new NoImageGeneratedError({ responses });
|
|
339
|
+
throw new NoImageGeneratedError({ calls, responses });
|
|
299
340
|
}
|
|
300
341
|
|
|
301
342
|
return new DefaultGenerateImageResult({
|