ai 7.0.93 → 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.
@@ -1,7 +1,9 @@
1
1
  import {
2
2
  getErrorMessage,
3
+ type LanguageModelV4Content,
3
4
  type LanguageModelV4Prompt,
4
5
  type LanguageModelV4StreamPart,
6
+ type LanguageModelV4ToolChoice,
5
7
  type SharedV4Headers,
6
8
  } from '@ai-sdk/provider';
7
9
  import {
@@ -40,6 +42,7 @@ import {
40
42
  import type { DownloadFunction } from '../util/download/download-function';
41
43
  import { notify } from '../util/notify';
42
44
  import { now as originalNow } from '../util/now';
45
+ import { ToolChoiceViolationError } from '../error';
43
46
  import { calculateTokensPerSecond } from './calculate-tokens-per-second';
44
47
  import type { ContentPart } from './content-part';
45
48
  import { DefaultGeneratedFileWithType } from './generated-file';
@@ -372,6 +375,7 @@ export async function streamLanguageModelCall<
372
375
  callId: effectiveCallId,
373
376
  provider: resolvedModel.provider,
374
377
  modelId: resolvedModel.modelId,
378
+ toolChoice: stepToolChoice,
375
379
  generateId,
376
380
  now,
377
381
  callStartTimestampMs,
@@ -398,6 +402,7 @@ function createLanguageModelV4StreamPartToLanguageModelStreamPartTransform<
398
402
  callId,
399
403
  provider,
400
404
  modelId,
405
+ toolChoice,
401
406
  generateId,
402
407
  now,
403
408
  callStartTimestampMs,
@@ -411,6 +416,7 @@ function createLanguageModelV4StreamPartToLanguageModelStreamPartTransform<
411
416
  callId: string;
412
417
  provider: string;
413
418
  modelId: string;
419
+ toolChoice: LanguageModelV4ToolChoice;
414
420
  generateId: IdGenerator;
415
421
  now: () => number;
416
422
  callStartTimestampMs: number;
@@ -420,8 +426,11 @@ function createLanguageModelV4StreamPartToLanguageModelStreamPartTransform<
420
426
  // keep track of tool inputs for provider-side tool results
421
427
  const toolCallsByToolCallId = new Map<string, TypedToolCall<TOOLS>>();
422
428
  const modelCallContent: Array<ContentPart<TOOLS>> = [];
429
+ const rawModelCallContent: Array<LanguageModelV4Content> = [];
423
430
  const textPartIndexes = new Map<string, number>();
424
431
  const reasoningPartIndexes = new Map<string, number>();
432
+ const rawTextPartIndexes = new Map<string, number>();
433
+ const rawReasoningPartIndexes = new Map<string, number>();
425
434
  let responseId = generateId();
426
435
  let responseModelId = modelId;
427
436
  let timeToFirstOutputMs: number | undefined;
@@ -458,7 +467,9 @@ function createLanguageModelV4StreamPartToLanguageModelStreamPartTransform<
458
467
  case 'text-start':
459
468
  upsertTextContentPart({
460
469
  content: modelCallContent,
470
+ rawContent: rawModelCallContent,
461
471
  partIndexes: textPartIndexes,
472
+ rawPartIndexes: rawTextPartIndexes,
462
473
  id: chunk.id,
463
474
  type: 'text',
464
475
  providerMetadata: chunk.providerMetadata,
@@ -469,7 +480,9 @@ function createLanguageModelV4StreamPartToLanguageModelStreamPartTransform<
469
480
  case 'text-delta':
470
481
  upsertTextContentPart({
471
482
  content: modelCallContent,
483
+ rawContent: rawModelCallContent,
472
484
  partIndexes: textPartIndexes,
485
+ rawPartIndexes: rawTextPartIndexes,
473
486
  id: chunk.id,
474
487
  type: 'text',
475
488
  textDelta: chunk.delta,
@@ -486,19 +499,24 @@ function createLanguageModelV4StreamPartToLanguageModelStreamPartTransform<
486
499
  case 'text-end':
487
500
  upsertTextContentPart({
488
501
  content: modelCallContent,
502
+ rawContent: rawModelCallContent,
489
503
  partIndexes: textPartIndexes,
504
+ rawPartIndexes: rawTextPartIndexes,
490
505
  id: chunk.id,
491
506
  type: 'text',
492
507
  providerMetadata: chunk.providerMetadata,
493
508
  });
494
509
  textPartIndexes.delete(chunk.id);
510
+ rawTextPartIndexes.delete(chunk.id);
495
511
  controller.enqueue(chunk);
496
512
  break;
497
513
 
498
514
  case 'reasoning-start':
499
515
  upsertTextContentPart({
500
516
  content: modelCallContent,
517
+ rawContent: rawModelCallContent,
501
518
  partIndexes: reasoningPartIndexes,
519
+ rawPartIndexes: rawReasoningPartIndexes,
502
520
  id: chunk.id,
503
521
  type: 'reasoning',
504
522
  providerMetadata: chunk.providerMetadata,
@@ -509,7 +527,9 @@ function createLanguageModelV4StreamPartToLanguageModelStreamPartTransform<
509
527
  case 'reasoning-delta':
510
528
  upsertTextContentPart({
511
529
  content: modelCallContent,
530
+ rawContent: rawModelCallContent,
512
531
  partIndexes: reasoningPartIndexes,
532
+ rawPartIndexes: rawReasoningPartIndexes,
513
533
  id: chunk.id,
514
534
  type: 'reasoning',
515
535
  textDelta: chunk.delta,
@@ -526,12 +546,15 @@ function createLanguageModelV4StreamPartToLanguageModelStreamPartTransform<
526
546
  case 'reasoning-end':
527
547
  upsertTextContentPart({
528
548
  content: modelCallContent,
549
+ rawContent: rawModelCallContent,
529
550
  partIndexes: reasoningPartIndexes,
551
+ rawPartIndexes: rawReasoningPartIndexes,
530
552
  id: chunk.id,
531
553
  type: 'reasoning',
532
554
  providerMetadata: chunk.providerMetadata,
533
555
  });
534
556
  reasoningPartIndexes.delete(chunk.id);
557
+ rawReasoningPartIndexes.delete(chunk.id);
535
558
  controller.enqueue(chunk);
536
559
  break;
537
560
 
@@ -552,6 +575,7 @@ function createLanguageModelV4StreamPartToLanguageModelStreamPartTransform<
552
575
  ? { providerMetadata: chunk.providerMetadata }
553
576
  : {}),
554
577
  });
578
+ rawModelCallContent.push(chunk);
555
579
 
556
580
  controller.enqueue({
557
581
  type: chunk.type,
@@ -612,6 +636,9 @@ function createLanguageModelV4StreamPartToLanguageModelStreamPartTransform<
612
636
  callbacks: onLanguageModelCallEnd,
613
637
  });
614
638
 
639
+ // Preserve the completed model call's usage, metadata, and
640
+ // performance even when response validation below surfaces a
641
+ // semantic error.
615
642
  controller.enqueue({
616
643
  type: 'model-call-end',
617
644
  finishReason: chunk.finishReason.unified,
@@ -620,10 +647,38 @@ function createLanguageModelV4StreamPartToLanguageModelStreamPartTransform<
620
647
  providerMetadata: chunk.providerMetadata,
621
648
  performance,
622
649
  });
650
+
651
+ const enforcedToolChoice =
652
+ toolChoice.type === 'required' || toolChoice.type === 'tool'
653
+ ? toolChoice
654
+ : undefined;
655
+
656
+ if (
657
+ enforcedToolChoice != null &&
658
+ ![...toolCallsByToolCallId.values()].some(
659
+ toolCall =>
660
+ enforcedToolChoice.type === 'required' ||
661
+ toolCall.toolName === enforcedToolChoice.toolName,
662
+ )
663
+ ) {
664
+ controller.enqueue({
665
+ type: 'error',
666
+ error: new ToolChoiceViolationError({
667
+ toolChoice: enforcedToolChoice,
668
+ finishReason: chunk.finishReason.unified,
669
+ provider,
670
+ modelId,
671
+ content: rawModelCallContent,
672
+ }),
673
+ });
674
+ break;
675
+ }
623
676
  break;
624
677
  }
625
678
 
626
679
  case 'tool-call': {
680
+ rawModelCallContent.push(chunk);
681
+
627
682
  try {
628
683
  const toolCall = await parseToolCall({
629
684
  toolCall: chunk,
@@ -663,6 +718,8 @@ function createLanguageModelV4StreamPartToLanguageModelStreamPartTransform<
663
718
  }
664
719
 
665
720
  case 'tool-approval-request': {
721
+ rawModelCallContent.push(chunk);
722
+
666
723
  const toolCall = toolCallsByToolCallId.get(chunk.toolCallId);
667
724
 
668
725
  if (toolCall == null) {
@@ -688,6 +745,8 @@ function createLanguageModelV4StreamPartToLanguageModelStreamPartTransform<
688
745
  }
689
746
 
690
747
  case 'tool-result': {
748
+ rawModelCallContent.push(chunk);
749
+
691
750
  const toolName = chunk.toolName as keyof TOOLS & string;
692
751
  const toolCall = toolCallsByToolCallId.get(chunk.toolCallId);
693
752
 
@@ -765,6 +824,7 @@ function createLanguageModelV4StreamPartToLanguageModelStreamPartTransform<
765
824
  default:
766
825
  if (chunk.type === 'custom' || chunk.type === 'source') {
767
826
  modelCallContent.push(chunk);
827
+ rawModelCallContent.push(chunk);
768
828
  }
769
829
 
770
830
  controller.enqueue(chunk);
@@ -818,14 +878,18 @@ function calculateNearestRankPercentile(
818
878
  */
819
879
  function upsertTextContentPart<TOOLS extends ToolSet>({
820
880
  content,
881
+ rawContent,
821
882
  partIndexes,
883
+ rawPartIndexes,
822
884
  id,
823
885
  type,
824
886
  textDelta,
825
887
  providerMetadata,
826
888
  }: {
827
889
  content: Array<ContentPart<TOOLS>>;
890
+ rawContent: Array<LanguageModelV4Content>;
828
891
  partIndexes: Map<string, number>;
892
+ rawPartIndexes: Map<string, number>;
829
893
  id: string;
830
894
  type: 'text' | 'reasoning';
831
895
  textDelta?: string;
@@ -843,16 +907,34 @@ function upsertTextContentPart<TOOLS extends ToolSet>({
843
907
  partIndexes.set(id, partIndex);
844
908
  }
845
909
 
910
+ let rawPartIndex = rawPartIndexes.get(id);
911
+
912
+ if (rawPartIndex == null) {
913
+ rawPartIndex =
914
+ rawContent.push({
915
+ type,
916
+ text: '',
917
+ ...(providerMetadata != null ? { providerMetadata } : {}),
918
+ }) - 1;
919
+ rawPartIndexes.set(id, rawPartIndex);
920
+ }
921
+
846
922
  const part = content[partIndex] as {
847
923
  text: string;
848
924
  providerMetadata?: ProviderMetadata;
849
925
  };
926
+ const rawPart = rawContent[rawPartIndex] as {
927
+ text: string;
928
+ providerMetadata?: ProviderMetadata;
929
+ };
850
930
 
851
931
  if (textDelta != null) {
852
932
  part.text += textDelta;
933
+ rawPart.text += textDelta;
853
934
  }
854
935
 
855
936
  if (providerMetadata != null) {
856
937
  part.providerMetadata = providerMetadata;
938
+ rawPart.providerMetadata = providerMetadata;
857
939
  }
858
940
  }
@@ -22,7 +22,7 @@ import {
22
22
  type ToolSet,
23
23
  } from '@ai-sdk/provider-utils';
24
24
  import type { ServerResponse } from 'node:http';
25
- import { NoOutputGeneratedError } from '../error';
25
+ import { NoOutputGeneratedError, ToolChoiceViolationError } from '../error';
26
26
  import { logWarnings } from '../logger/log-warnings';
27
27
  import { resolveLanguageModel } from '../model/resolve-model';
28
28
  import { cloneModelMessages } from '../prompt/clone-model-message';
@@ -2573,6 +2573,8 @@ class DefaultStreamTextResult<
2573
2573
  callbacks: onChunk,
2574
2574
  });
2575
2575
  const error = wrapGatewayError(value.error);
2576
+ const isToolChoiceViolation =
2577
+ ToolChoiceViolationError.isInstance(error);
2576
2578
  let onErrorResult: unknown;
2577
2579
  try {
2578
2580
  onErrorResult = await onError({ error });
@@ -2584,8 +2586,10 @@ class DefaultStreamTextResult<
2584
2586
  'retry' in onErrorResult &&
2585
2587
  onErrorResult.retry === true;
2586
2588
  const automaticRetry =
2589
+ !isToolChoiceViolation &&
2587
2590
  automaticStreamRetryCount < streamRetries;
2588
2591
  const callbackRetry =
2592
+ !isToolChoiceViolation &&
2589
2593
  !automaticRetry &&
2590
2594
  callbackRequestedRetry &&
2591
2595
  callbackStreamRetryCount < 1;
@@ -1,5 +1,8 @@
1
1
  import { InvalidArgumentError } from '../error/invalid-argument-error';
2
- import type { RetryFunction } from '@ai-sdk/provider-utils';
2
+ import type {
3
+ RetryFunction,
4
+ ShouldRetryFunction,
5
+ } from '@ai-sdk/provider-utils';
3
6
  import { retryWithExponentialBackoffRespectingRetryHeaders } from '../util/retry-with-exponential-backoff';
4
7
  /**
5
8
  * Validate and prepare retries.
@@ -7,11 +10,13 @@ import { retryWithExponentialBackoffRespectingRetryHeaders } from '../util/retry
7
10
  export function prepareRetries({
8
11
  maxRetries,
9
12
  abortSignal,
13
+ additionalRetryableError,
10
14
  parameter = 'maxRetries',
11
15
  defaultMaxRetries = 2,
12
16
  }: {
13
17
  maxRetries: number | undefined;
14
18
  abortSignal: AbortSignal | undefined;
19
+ additionalRetryableError?: ShouldRetryFunction;
15
20
  parameter?: string;
16
21
  defaultMaxRetries?: number;
17
22
  }): {
@@ -43,6 +48,7 @@ export function prepareRetries({
43
48
  retry: retryWithExponentialBackoffRespectingRetryHeaders({
44
49
  maxRetries: maxRetriesResult,
45
50
  abortSignal,
51
+ additionalRetryableError,
46
52
  }),
47
53
  };
48
54
  }
@@ -3,6 +3,7 @@ import { GatewayError } from '@ai-sdk/gateway';
3
3
  import {
4
4
  retryWithExponentialBackoff,
5
5
  type RetryFunction,
6
+ type ShouldRetryFunction,
6
7
  } from '@ai-sdk/provider-utils';
7
8
  import { RetryError } from './retry-error';
8
9
 
@@ -66,21 +67,25 @@ export const retryWithExponentialBackoffRespectingRetryHeaders = ({
66
67
  initialDelayInMs = 2000,
67
68
  backoffFactor = 2,
68
69
  abortSignal,
70
+ additionalRetryableError,
69
71
  }: {
70
72
  maxRetries?: number;
71
73
  initialDelayInMs?: number;
72
74
  backoffFactor?: number;
73
75
  abortSignal?: AbortSignal;
76
+ additionalRetryableError?: ShouldRetryFunction;
74
77
  } = {}): RetryFunction =>
75
78
  retryWithExponentialBackoff({
76
79
  maxRetries,
77
80
  initialDelayInMs,
78
81
  backoffFactor,
79
82
  abortSignal,
80
- shouldRetry: error =>
81
- error instanceof Error &&
82
- ((APICallError.isInstance(error) && error.isRetryable === true) ||
83
- (GatewayError.isInstance(error) && error.isRetryable === true)),
83
+ shouldRetry: async error =>
84
+ (error instanceof Error &&
85
+ ((APICallError.isInstance(error) && error.isRetryable === true) ||
86
+ (GatewayError.isInstance(error) && error.isRetryable === true))) ||
87
+ (additionalRetryableError != null &&
88
+ (await additionalRetryableError(error))),
84
89
  getDelayInMs: ({ error, exponentialBackoffDelay }) =>
85
90
  getRetryDelayInMs({
86
91
  error: error as APICallError | GatewayError,