@f5-sales-demo/pi-agent-core 22.7.1 → 22.7.2

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/package.json CHANGED
@@ -1,7 +1,7 @@
1
1
  {
2
2
  "type": "module",
3
3
  "name": "@f5-sales-demo/pi-agent-core",
4
- "version": "22.7.1",
4
+ "version": "22.7.2",
5
5
  "description": "General-purpose agent with transport abstraction, state management, and attachment support",
6
6
  "homepage": "https://github.com/f5-sales-demo/xcsh",
7
7
  "author": "Can Boluk",
@@ -35,8 +35,8 @@
35
35
  "fmt": "biome format --write ."
36
36
  },
37
37
  "dependencies": {
38
- "@f5-sales-demo/pi-ai": "22.7.1",
39
- "@f5-sales-demo/pi-utils": "22.7.1"
38
+ "@f5-sales-demo/pi-ai": "22.7.2",
39
+ "@f5-sales-demo/pi-utils": "22.7.2"
40
40
  },
41
41
  "devDependencies": {
42
42
  "@sinclair/typebox": "0.34.52",
package/src/agent-loop.ts CHANGED
@@ -188,10 +188,16 @@ async function bufferAssistantResponse(
188
188
  response: Awaited<ReturnType<StreamFn>>,
189
189
  model: AgentLoopConfig["model"],
190
190
  signal: AbortSignal | undefined,
191
+ onConfirmed?: (
192
+ message: AssistantMessage,
193
+ call: Extract<AssistantMessage["content"][number], { type: "toolCall" }>,
194
+ ) => Promise<void>,
195
+ requiredName?: string,
191
196
  ): Promise<BufferedAssistantResponse> {
192
197
  const events: BufferedAssistantResponse["events"] = [];
193
198
  let partialMessage: AssistantMessage | null = null;
194
199
  let addedPartial = false;
200
+ const completed = new Set<number>();
195
201
 
196
202
  for await (const event of response) {
197
203
  if (signal?.aborted) {
@@ -224,11 +230,25 @@ async function bufferAssistantResponse(
224
230
  case "server_tool_end":
225
231
  if (partialMessage) {
226
232
  partialMessage = event.partial;
233
+ if (event.type === "toolcall_end" || event.type === "thinking_end" || event.type === "text_end")
234
+ completed.add(event.contentIndex);
227
235
  events.push({
228
236
  type: "message_update",
229
237
  assistantMessageEvent: event,
230
238
  message: { ...partialMessage } as AssistantMessage,
231
239
  });
240
+ if (
241
+ event.type === "toolcall_end" &&
242
+ onConfirmed &&
243
+ event.toolCall.name === requiredName &&
244
+ event.partial.content.filter(c => c.type === "toolCall").length === 1
245
+ ) {
246
+ await onConfirmed(
247
+ { ...event.partial, content: event.partial.content.filter((_c, i) => completed.has(i)) },
248
+ event.toolCall,
249
+ );
250
+ onConfirmed = undefined;
251
+ }
232
252
  }
233
253
  break;
234
254
  case "done":
@@ -325,11 +345,14 @@ async function runLoop(
325
345
  message.model === config.model.id
326
346
  )
327
347
  for (const call of message.content) if (call.type === "toolCall") priorCalls.set(call.id, toolKey(call));
328
- if (message.role === "toolResult" && !message.isError) {
348
+ if (message.role === "toolResult") {
329
349
  const key = priorCalls.get(message.toolCallId);
330
350
  if (key) completedTools.set(key, message);
331
351
  }
332
352
  }
353
+ let recoveryAttempts = 0;
354
+ let recoveryToolChoiceServed = false;
355
+ let recoveryToolChoice: { value: AgentLoopConfig["toolChoice"] } | undefined;
333
356
  let firstTurn = true;
334
357
  // Check for steering messages at start (user may have typed while waiting)
335
358
  let pendingMessages: AgentMessage[] = (await config.getSteeringMessages?.()) || [];
@@ -362,13 +385,17 @@ async function runLoop(
362
385
  await logger.ttftAttr("ttft.sync-context", () => config.syncContextBeforeModelCall!(currentContext));
363
386
  }
364
387
 
388
+ const selectedToolChoice = recoveryToolChoice
389
+ ? recoveryToolChoice.value
390
+ : (config.getToolChoice?.() ?? config.toolChoice);
391
+
365
392
  // Offer opt-in, plain user steering to the active provider without consuming other message kinds.
366
393
  const live =
367
394
  (config.model.compat as import("@f5-sales-demo/pi-ai").OpenAIResponsesCompat | undefined)
368
395
  ?.supportsWebSocketSteering &&
369
396
  config.waitForSteeringMessages &&
370
397
  config.getSteeringMessages &&
371
- !exactToolName(config.getToolChoice?.() ?? config.toolChoice)
398
+ !exactToolName(selectedToolChoice)
372
399
  ? new LiveSteeringChannel({
373
400
  wait: config.waitForSteeringMessages,
374
401
  take: () => config.getSteeringMessages!(),
@@ -379,15 +406,89 @@ async function runLoop(
379
406
  },
380
407
  })
381
408
  : undefined;
382
- const message = await streamAssistantResponse(
383
- currentContext,
384
- newMessages,
385
- live ? { ...config, liveSteering: live } : config,
409
+ const scheduler = createToolScheduler(
410
+ currentContext.tools,
386
411
  signal,
387
412
  stream,
388
- streamFn,
413
+ config.getSteeringMessages,
414
+ config.interruptMode,
415
+ config.getToolContext,
416
+ config.transformToolCallArguments,
417
+ config.intentTracing,
418
+ completedTools,
419
+ toolKey,
389
420
  );
390
- newMessages.push(message);
421
+ const streamingDispatch = config.model.api === "openai-codex-responses";
422
+ const responseController = new AbortController();
423
+ const responseSignal = signal
424
+ ? AbortSignal.any([signal, responseController.signal])
425
+ : responseController.signal;
426
+ const responseStartIndex = currentContext.messages.length;
427
+ let message: AssistantMessage;
428
+ let lastCheckpoint: AssistantMessage | undefined;
429
+ try {
430
+ message = await streamAssistantResponse(
431
+ currentContext,
432
+ newMessages,
433
+ {
434
+ ...config,
435
+ ...(live ? { liveSteering: live } : {}),
436
+ onAssistantCheckpoint: async checkpoint => {
437
+ await config.onAssistantCheckpoint?.(checkpoint);
438
+ lastCheckpoint = checkpoint;
439
+ },
440
+ getToolChoice: undefined,
441
+ toolChoice: selectedToolChoice,
442
+ ...(streamingDispatch ? { maxRetries: Math.max(0, 2 - recoveryAttempts) } : {}),
443
+ },
444
+ responseSignal,
445
+ stream,
446
+ streamFn,
447
+ streamingDispatch ? (partial, call) => scheduler.admit(partial, [call]) : undefined,
448
+ );
449
+ } catch (error) {
450
+ responseController.abort();
451
+ const prior =
452
+ lastCheckpoint ?? currentContext.messages.slice(responseStartIndex).find(m => m.role === "assistant");
453
+ message =
454
+ prior?.role === "assistant"
455
+ ? {
456
+ ...prior,
457
+ stopReason: signal?.aborted ? "aborted" : "error",
458
+ errorMessage: error instanceof Error ? error.message : String(error),
459
+ }
460
+ : {
461
+ ...sanitizedForcedToolMessage(config.model, signal?.aborted ? "aborted" : "error"),
462
+ errorMessage: String(error),
463
+ };
464
+ const ownedIndex = currentContext.messages.findIndex(
465
+ (m, index) => index >= responseStartIndex && m.role === "assistant" && m.timestamp === message.timestamp,
466
+ );
467
+ if (ownedIndex >= 0) currentContext.messages[ownedIndex] = message;
468
+ else currentContext.messages.push(message);
469
+ stream.push({ type: "message_end", message });
470
+ }
471
+ newMessages.push(retainedAssistant(message));
472
+ if (message.stopReason !== "error" && message.stopReason !== "aborted") {
473
+ recoveryAttempts = 0;
474
+ recoveryToolChoice = undefined;
475
+ recoveryToolChoiceServed = false;
476
+ const required = exactToolName(selectedToolChoice);
477
+ scheduler.admit(
478
+ message,
479
+ message.content.filter(
480
+ (c): c is Extract<AssistantMessage["content"][number], { type: "toolCall" }> =>
481
+ c.type === "toolCall" && (!required || c.name === required),
482
+ ),
483
+ );
484
+ }
485
+ const executionResult = await scheduler.drain();
486
+ const dispatchedResults = executionResult.toolResults;
487
+ for (const result of dispatchedResults) {
488
+ if (currentContext.messages.includes(result)) continue;
489
+ currentContext.messages.push(result);
490
+ newMessages.push(result);
491
+ }
391
492
  if (live) {
392
493
  if (message.stopReason === "error" || message.stopReason === "aborted")
393
494
  config.restoreSteeringMessages?.([...live.accepted, ...live.deferred]);
@@ -401,21 +502,57 @@ async function runLoop(
401
502
  config.restoreSteeringMessages?.(live.deferred);
402
503
  }
403
504
  }
404
- let steeringMessagesFromExecution: AgentMessage[] | undefined;
405
-
505
+ const steeringMessagesFromExecution = executionResult.steeringMessages;
506
+
507
+ if (streamingDispatch && message.interruption && message.stopReason === "error" && !signal?.aborted) {
508
+ recoveryToolChoiceServed ||= !!exactToolName(selectedToolChoice) && dispatchedResults.length > 0;
509
+ if (recoveryToolChoiceServed) message.interruption.toolChoiceServed = true;
510
+ recoveryAttempts += message.interruption.providerRetriesConsumed;
511
+ if (recoveryAttempts < 2) {
512
+ recoveryAttempts += 1;
513
+ try {
514
+ await recoveryBackoff(500 * recoveryAttempts, signal);
515
+ } catch {
516
+ message.stopReason = "aborted";
517
+ const retained = retainedAssistant(message);
518
+ const index = newMessages.findIndex(m => m.role === "assistant" && m.timestamp === message.timestamp);
519
+ if (index >= 0) newMessages[index] = retained;
520
+ const contextIndex = currentContext.messages.findIndex(
521
+ m => m.role === "assistant" && m.timestamp === message.timestamp,
522
+ );
523
+ if (contextIndex >= 0) currentContext.messages[contextIndex] = retained;
524
+ stream.push({ type: "message_end", message: retained });
525
+ }
526
+ if (!signal?.aborted) {
527
+ pendingMessages = steeringMessagesFromExecution ?? ((await config.getSteeringMessages?.()) || []);
528
+ hasMoreToolCalls = true;
529
+ recoveryToolChoice = {
530
+ value:
531
+ exactToolName(selectedToolChoice) && dispatchedResults.length > 0 ? "none" : selectedToolChoice,
532
+ };
533
+ continue;
534
+ }
535
+ }
536
+ }
406
537
  if (message.stopReason === "error" || message.stopReason === "aborted") {
407
538
  // Create placeholder tool results for any tool calls in the aborted message
408
539
  // This maintains the tool_use/tool_result pairing that the API requires
409
540
  type ToolCallContent = Extract<AssistantMessage["content"][number], { type: "toolCall" }>;
410
541
  const toolCalls = message.content.filter((c): c is ToolCallContent => c.type === "toolCall");
411
- const toolResults: ToolResultMessage[] = [];
542
+ const toolResults: ToolResultMessage[] = [...dispatchedResults];
412
543
  for (const toolCall of toolCalls) {
544
+ if (scheduler.has(toolCall) || message.interruption) continue;
413
545
  const result = createAbortedToolResult(toolCall, stream, message.stopReason, message.errorMessage);
414
546
  currentContext.messages.push(result);
415
547
  newMessages.push(result);
416
548
  toolResults.push(result);
417
549
  }
418
- stream.push({ type: "turn_end", message, toolResults });
550
+ stream.push({
551
+ type: "turn_end",
552
+ message,
553
+ toolResults,
554
+ ...(recoveryToolChoiceServed ? { toolChoiceServed: true } : {}),
555
+ });
419
556
  stream.push({ type: "agent_end", messages: newMessages });
420
557
  stream.end(newMessages);
421
558
  return;
@@ -425,33 +562,14 @@ async function runLoop(
425
562
  const toolCalls = message.content.filter(c => c.type === "toolCall");
426
563
  hasMoreToolCalls = toolCalls.length > 0 || Boolean(live?.accepted.length);
427
564
 
428
- const toolResults: ToolResultMessage[] = [];
429
- if (toolCalls.length > 0) {
430
- const executionResult = await executeToolCalls(
431
- currentContext.tools,
432
- message,
433
- signal,
434
- stream,
435
- config.getSteeringMessages,
436
- config.interruptMode,
437
- config.getToolContext,
438
- config.transformToolCallArguments,
439
- config.intentTracing,
440
- completedTools,
441
- toolKey,
442
- );
443
-
444
- toolResults.push(...executionResult.toolResults);
445
- steeringMessagesFromExecution = executionResult.steeringMessages;
446
-
447
- for (const result of toolResults) {
448
- if (currentContext.messages.includes(result)) continue;
449
- currentContext.messages.push(result);
450
- newMessages.push(result);
451
- }
452
- }
565
+ const toolResults = dispatchedResults;
453
566
 
454
- stream.push({ type: "turn_end", message, toolResults });
567
+ stream.push({
568
+ type: "turn_end",
569
+ message,
570
+ toolResults,
571
+ ...(recoveryToolChoiceServed ? { toolChoiceServed: true } : {}),
572
+ });
455
573
 
456
574
  pendingMessages = steeringMessagesFromExecution ?? ((await config.getSteeringMessages?.()) || []);
457
575
  }
@@ -483,6 +601,10 @@ async function streamAssistantResponse(
483
601
  signal: AbortSignal | undefined,
484
602
  stream: EventStream<AgentEvent, AgentMessage[]>,
485
603
  streamFn?: StreamFn,
604
+ onCompletedCall?: (
605
+ message: AssistantMessage,
606
+ call: Extract<AssistantMessage["content"][number], { type: "toolCall" }>,
607
+ ) => void,
486
608
  ): Promise<AssistantMessage> {
487
609
  // Apply context transform if configured (AgentMessage[] → AgentMessage[])
488
610
  let messages = context.messages;
@@ -533,17 +655,68 @@ async function streamAssistantResponse(
533
655
  signal,
534
656
  }),
535
657
  );
536
- const buffered = await bufferAssistantResponse(response, config.model, signal);
658
+ let confirmed: AssistantMessage | undefined;
659
+ const buffered = await bufferAssistantResponse(
660
+ response,
661
+ config.model,
662
+ signal,
663
+ onCompletedCall
664
+ ? async (partial, call) => {
665
+ confirmed = structuredClone({
666
+ ...partial,
667
+ content: partial.content.filter(c => c.type !== "toolCall" || c.id === call.id),
668
+ });
669
+ await config.onAssistantCheckpoint?.(confirmed);
670
+ stream.push({ type: "assistant_checkpoint", message: confirmed });
671
+ onCompletedCall(confirmed, call);
672
+ }
673
+ : undefined,
674
+ requiredToolName,
675
+ );
676
+ if (confirmed && (buffered.message.stopReason === "error" || buffered.message.stopReason === "aborted")) {
677
+ const failure = {
678
+ ...buffered.message,
679
+ content: confirmed.content,
680
+ providerPayload: confirmed.providerPayload,
681
+ };
682
+ if (failure.interruption)
683
+ failure.interruption = {
684
+ ...failure.interruption,
685
+ completedContentIndices: failure.content.map((_c, i) => i),
686
+ };
687
+ context.messages.push(retainedAssistant(failure));
688
+ stream.push({
689
+ type: "message_end",
690
+ message: quietAssistant(failure, !!onCompletedCall, config.maxRetries ?? 2, signal),
691
+ });
692
+ return failure;
693
+ }
537
694
  if (buffered.message.stopReason === "aborted") {
538
695
  context.messages.push(buffered.message);
539
696
  for (const event of buffered.events) stream.push(event);
540
697
  return buffered.message;
541
698
  }
542
- if (buffered.message.stopReason === "error") break;
699
+ if (buffered.message.stopReason === "error") {
700
+ if (onCompletedCall && buffered.message.interruption) {
701
+ context.messages.push(retainedAssistant(buffered.message));
702
+ stream.push({
703
+ type: "message_end",
704
+ message: quietAssistant(buffered.message, !!onCompletedCall, config.maxRetries ?? 2, signal),
705
+ });
706
+ return buffered.message;
707
+ }
708
+ break;
709
+ }
543
710
 
544
711
  const toolCalls = buffered.message.content.filter(content => content.type === "toolCall");
545
712
  const exactInvocation = toolCalls.length === 1 && toolCalls[0]?.name === requiredToolName;
546
713
  if (exactInvocation) {
714
+ if (onCompletedCall && !confirmed) {
715
+ const checkpoint = structuredClone(buffered.message);
716
+ await config.onAssistantCheckpoint?.(checkpoint);
717
+ stream.push({ type: "assistant_checkpoint", message: checkpoint });
718
+ onCompletedCall(checkpoint, toolCalls[0]);
719
+ }
547
720
  context.messages.push(buffered.message);
548
721
  for (const event of buffered.events) {
549
722
  if (event.type === "message_update") {
@@ -554,6 +727,19 @@ async function streamAssistantResponse(
554
727
  return buffered.message;
555
728
  }
556
729
 
730
+ if (confirmed) {
731
+ const failure = {
732
+ ...confirmed,
733
+ stopReason: "error" as const,
734
+ errorMessage: "Required tool invocation failed after a completed call.",
735
+ };
736
+ context.messages.push(failure);
737
+ stream.push({
738
+ type: "message_end",
739
+ message: quietAssistant(failure, !!onCompletedCall, config.maxRetries ?? 2, signal),
740
+ });
741
+ return failure;
742
+ }
557
743
  if (attempt === 0 && buffered.message.stopReason === "length" && toolCalls.length === 0) continue;
558
744
  break;
559
745
  }
@@ -576,7 +762,10 @@ async function streamAssistantResponse(
576
762
 
577
763
  let partialMessage: AssistantMessage | null = null;
578
764
  let addedPartial = false;
765
+ let assistantIndex = -1;
766
+ const completedIndices = new Set<number>();
579
767
  let firstDeltaMarked = false;
768
+ let namedCallAdmitted = false;
580
769
 
581
770
  for await (const event of response) {
582
771
  if (!firstDeltaMarked && event.type === "text_delta") {
@@ -607,7 +796,7 @@ async function streamAssistantResponse(
607
796
  timestamp: Date.now(),
608
797
  };
609
798
  if (addedPartial) {
610
- context.messages[context.messages.length - 1] = abortedMessage;
799
+ context.messages[assistantIndex] = abortedMessage;
611
800
  } else {
612
801
  context.messages.push(abortedMessage);
613
802
  stream.push({ type: "message_start", message: { ...abortedMessage } });
@@ -619,6 +808,7 @@ async function streamAssistantResponse(
619
808
  switch (event.type) {
620
809
  case "start":
621
810
  partialMessage = event.partial;
811
+ assistantIndex = context.messages.length;
622
812
  context.messages.push(partialMessage);
623
813
  addedPartial = true;
624
814
  stream.push({ type: "message_start", message: { ...partialMessage } });
@@ -640,8 +830,28 @@ async function streamAssistantResponse(
640
830
  case "server_tool_end":
641
831
  if (partialMessage) {
642
832
  partialMessage = event.partial;
643
- context.messages[context.messages.length - 1] = partialMessage;
833
+ context.messages[assistantIndex] = partialMessage;
644
834
  config.onAssistantMessageEvent?.(partialMessage, event);
835
+ if (
836
+ onCompletedCall &&
837
+ (event.type === "toolcall_end" || event.type === "thinking_end" || event.type === "text_end")
838
+ ) {
839
+ completedIndices.add(event.contentIndex);
840
+ const checkpoint = structuredClone({
841
+ ...partialMessage,
842
+ content: partialMessage.content.filter((_c, i) => completedIndices.has(i)),
843
+ });
844
+ await config.onAssistantCheckpoint?.(checkpoint);
845
+ stream.push({ type: "assistant_checkpoint", message: checkpoint });
846
+ if (
847
+ event.type === "toolcall_end" &&
848
+ !signal?.aborted &&
849
+ (!requiredToolName || (event.toolCall.name === requiredToolName && !namedCallAdmitted))
850
+ ) {
851
+ namedCallAdmitted = true;
852
+ onCompletedCall(checkpoint, event.toolCall);
853
+ }
854
+ }
645
855
  if (signal?.aborted) {
646
856
  continue;
647
857
  }
@@ -657,14 +867,17 @@ async function streamAssistantResponse(
657
867
  case "error": {
658
868
  const finalMessage = await response.result();
659
869
  if (addedPartial) {
660
- context.messages[context.messages.length - 1] = finalMessage;
870
+ context.messages[assistantIndex] = retainedAssistant(finalMessage);
661
871
  } else {
662
- context.messages.push(finalMessage);
872
+ context.messages.push(retainedAssistant(finalMessage));
663
873
  }
664
874
  if (!addedPartial) {
665
875
  stream.push({ type: "message_start", message: { ...finalMessage } });
666
876
  }
667
- stream.push({ type: "message_end", message: finalMessage });
877
+ stream.push({
878
+ type: "message_end",
879
+ message: quietAssistant(finalMessage, !!onCompletedCall, config.maxRetries ?? 2, signal),
880
+ });
668
881
  return finalMessage;
669
882
  }
670
883
  }
@@ -676,9 +889,8 @@ async function streamAssistantResponse(
676
889
  /**
677
890
  * Execute tool calls from an assistant message.
678
891
  */
679
- async function executeToolCalls(
892
+ function createToolScheduler(
680
893
  tools: AgentTool<any>[] | undefined,
681
- assistantMessage: AssistantMessage,
682
894
  signal: AbortSignal | undefined,
683
895
  stream: EventStream<AgentEvent, AgentMessage[]>,
684
896
  getSteeringMessages?: AgentLoopConfig["getSteeringMessages"],
@@ -688,18 +900,28 @@ async function executeToolCalls(
688
900
  intentTracing?: AgentLoopConfig["intentTracing"],
689
901
  completedTools = new Map<string, ToolResultMessage>(),
690
902
  toolKey: (call: Extract<AssistantMessage["content"][number], { type: "toolCall" }>) => string = call => call.id,
691
- ): Promise<{ toolResults: ToolResultMessage[]; steeringMessages?: AgentMessage[] }> {
903
+ ) {
692
904
  type ToolCallContent = Extract<AssistantMessage["content"][number], { type: "toolCall" }>;
693
- const toolCalls = [
694
- ...new Map(
695
- assistantMessage.content
696
- .filter((c): c is ToolCallContent => c.type === "toolCall")
697
- .map(call => [toolKey(call), call]),
698
- ).values(),
699
- ];
905
+ const advertisedTools = tools?.map(tool => ({
906
+ ...tool,
907
+ name: tool.name,
908
+ label: tool.label,
909
+ description: tool.description,
910
+ concurrency: tool.concurrency,
911
+ nonAbortable: tool.nonAbortable,
912
+ lenientArgValidation: tool.lenientArgValidation,
913
+ executionKind: tool.executionKind,
914
+ parameters: structuredClone(tool.parameters),
915
+ execute: tool.execute.bind(tool),
916
+ getExecutionKind: tool.getExecutionKind?.bind(tool),
917
+ }));
918
+ const toolCalls: ToolCallContent[] = [];
919
+ const admitted = new Set<string>();
920
+ const toolCallInfos: Array<{ id: string; name: string }> = [];
921
+ let batchId = "";
922
+
700
923
  const emittedToolResults: ToolResultMessage[] = [];
701
- const toolCallInfos = toolCalls.map(call => ({ id: call.id, name: call.name }));
702
- const batchId = `${assistantMessage.timestamp ?? Date.now()}_${toolCalls[0]?.id ?? "batch"}`;
924
+
703
925
  const shouldInterruptImmediately = interruptMode !== "wait";
704
926
  const steeringAbortController = new AbortController();
705
927
  const toolSignal = signal
@@ -709,9 +931,9 @@ async function executeToolCalls(
709
931
  let steeringMessages: AgentMessage[] | undefined;
710
932
  let steeringCheck: Promise<void> | null = null;
711
933
 
712
- const records = toolCalls.map(toolCall => ({
934
+ const makeRecord = (toolCall: ToolCallContent) => ({
713
935
  toolCall,
714
- tool: tools?.find(t => t.name === toolCall.name),
936
+ tool: advertisedTools?.find(t => t.name === toolCall.name),
715
937
  args: toolCall.arguments as Record<string, unknown>,
716
938
  started: false,
717
939
  result: undefined as AgentToolResult<any> | undefined,
@@ -719,7 +941,8 @@ async function executeToolCalls(
719
941
  skipped: false,
720
942
  toolResultMessage: undefined as ToolResultMessage | undefined,
721
943
  resultEmitted: false,
722
- }));
944
+ });
945
+ const records: ReturnType<typeof makeRecord>[] = [];
723
946
 
724
947
  const checkSteering = async (): Promise<void> => {
725
948
  if (!shouldInterruptImmediately || !getSteeringMessages || interruptState.triggered) {
@@ -783,7 +1006,7 @@ async function executeToolCalls(
783
1006
  record.isError = isError;
784
1007
  record.toolResultMessage = toolResultMessage;
785
1008
  record.resultEmitted = true;
786
- if (!isError) completedTools.set(toolKey(toolCall), toolResultMessage);
1009
+ completedTools.set(toolKey(toolCall), toolResultMessage);
787
1010
  emittedToolResults.push(toolResultMessage);
788
1011
 
789
1012
  stream.push({ type: "message_start", message: toolResultMessage });
@@ -798,7 +1021,7 @@ async function executeToolCalls(
798
1021
  emittedToolResults.push(completed);
799
1022
  return;
800
1023
  }
801
- if (interruptState.triggered) {
1024
+ if (interruptState.triggered || signal?.aborted) {
802
1025
  record.skipped = true;
803
1026
  return;
804
1027
  }
@@ -888,30 +1111,49 @@ async function executeToolCalls(
888
1111
  let sharedTasks: Promise<void>[] = [];
889
1112
  const tasks: Promise<void>[] = [];
890
1113
 
891
- for (let index = 0; index < records.length; index++) {
892
- const record = records[index];
893
- const concurrency = record.tool?.concurrency ?? "shared";
894
- const start = concurrency === "exclusive" ? Promise.all([lastExclusive, ...sharedTasks]) : lastExclusive;
895
- const task = start.then(() => runTool(record, index));
896
- tasks.push(task);
897
- if (concurrency === "exclusive") {
898
- lastExclusive = task;
899
- sharedTasks = [];
900
- } else {
901
- sharedTasks.push(task);
1114
+ const admit = (
1115
+ assistant: AssistantMessage,
1116
+ calls = assistant.content.filter((c): c is ToolCallContent => c.type === "toolCall"),
1117
+ ) => {
1118
+ if (signal?.aborted || interruptState.triggered) return;
1119
+ batchId ||= `${assistant.timestamp}_${calls[0]?.id ?? "batch"}`;
1120
+ for (const original of calls) {
1121
+ const key = toolKey(original);
1122
+ if (admitted.has(key)) continue;
1123
+ admitted.add(key);
1124
+ const call = structuredClone(original);
1125
+ if (intentTracing) {
1126
+ const { intent } = extractIntent(call.arguments);
1127
+ if (intent) original.intent = call.intent = intent;
1128
+ }
1129
+ toolCalls.push(call);
1130
+ toolCallInfos.push({ id: call.id, name: call.name });
1131
+ const record = makeRecord(call);
1132
+ const index = records.length;
1133
+ records.push(record);
1134
+ const concurrency = record.tool?.concurrency ?? "shared";
1135
+ const start = concurrency === "exclusive" ? Promise.all([lastExclusive, ...sharedTasks]) : lastExclusive;
1136
+ const task = start.then(() => runTool(record, index));
1137
+ tasks.push(task);
1138
+ if (concurrency === "exclusive") {
1139
+ lastExclusive = task;
1140
+ sharedTasks = [];
1141
+ } else sharedTasks.push(task);
902
1142
  }
903
- }
904
-
905
- await Promise.allSettled(tasks);
1143
+ };
1144
+ const drain = async () => {
1145
+ await Promise.allSettled(tasks);
906
1146
 
907
- for (const record of records) {
908
- if (!record.toolResultMessage) {
909
- record.skipped = true;
910
- emitToolResult(record, createSkippedToolResult(), true);
1147
+ for (const record of records) {
1148
+ if (!record.toolResultMessage) {
1149
+ record.skipped = true;
1150
+ emitToolResult(record, createSkippedToolResult(), true);
1151
+ }
911
1152
  }
912
- }
913
1153
 
914
- return { toolResults: emittedToolResults, steeringMessages };
1154
+ return { toolResults: emittedToolResults, steeringMessages };
1155
+ };
1156
+ return { admit, drain, has: (call: ToolCallContent) => admitted.has(toolKey(call)) };
915
1157
  }
916
1158
 
917
1159
  /**
@@ -924,7 +1166,7 @@ function createAbortedToolResult(
924
1166
  reason: "aborted" | "error",
925
1167
  errorMessage?: string,
926
1168
  ): ToolResultMessage {
927
- const message = reason === "aborted" ? "Tool execution was aborted" : "Tool execution failed due to an error";
1169
+ const message = reason === "aborted" ? "Tool execution was aborted" : "Response interrupted before tool execution";
928
1170
  const result: AgentToolResult<any> = {
929
1171
  content: [{ type: "text", text: errorMessage ? `${message}: ${errorMessage}` : `${message}.` }],
930
1172
  details: {},
@@ -978,3 +1220,44 @@ function normalizeToolResult<T>(result: AgentToolResult<T>, isError: boolean): A
978
1220
  content: [{ type: "text", text: emptyToolSuccess.trim() }],
979
1221
  };
980
1222
  }
1223
+
1224
+ function retainedAssistant(message: AssistantMessage): AssistantMessage {
1225
+ if (!message.interruption) return message;
1226
+ const indices = new Set(message.interruption.completedContentIndices);
1227
+ const content = message.content.filter((_c, i) => indices.has(i));
1228
+ return {
1229
+ ...message,
1230
+ content,
1231
+ interruption: { ...message.interruption, completedContentIndices: content.map((_c, i) => i) },
1232
+ };
1233
+ }
1234
+
1235
+ async function recoveryBackoff(delay: number, signal?: AbortSignal): Promise<void> {
1236
+ signal?.throwIfAborted();
1237
+ const { promise, resolve, reject } = Promise.withResolvers<void>();
1238
+ const timer = setTimeout(resolve, delay);
1239
+ const abort = () => reject(signal?.reason);
1240
+ signal?.addEventListener("abort", abort, { once: true });
1241
+ try {
1242
+ await promise;
1243
+ signal?.throwIfAborted();
1244
+ } finally {
1245
+ clearTimeout(timer);
1246
+ signal?.removeEventListener("abort", abort);
1247
+ }
1248
+ }
1249
+
1250
+ function quietAssistant(
1251
+ message: AssistantMessage,
1252
+ streaming: boolean,
1253
+ allowance: number,
1254
+ signal?: AbortSignal,
1255
+ ): AssistantMessage {
1256
+ const recoverable =
1257
+ streaming &&
1258
+ message.stopReason === "error" &&
1259
+ message.interruption &&
1260
+ !signal?.aborted &&
1261
+ message.interruption.providerRetriesConsumed < allowance;
1262
+ return recoverable ? { ...retainedAssistant(message), stopReason: "toolUse", errorMessage: undefined } : message;
1263
+ }
package/src/agent.ts CHANGED
@@ -138,6 +138,7 @@ export interface AgentOptions {
138
138
  * Inspect assistant streaming events before they are emitted to subscribers.
139
139
  * Use this when abort decisions must happen before buffered events continue flowing.
140
140
  */
141
+ onAssistantCheckpoint?: (message: AssistantMessage) => Promise<void>;
141
142
  onAssistantMessageEvent?: (message: AssistantMessage, event: AssistantMessageEvent) => void;
142
143
  /**
143
144
  * Custom token budgets for thinking levels (token-based providers only).
@@ -288,6 +289,8 @@ export class Agent {
288
289
  #getToolChoice?: () => ToolChoice | undefined;
289
290
  #onPayload?: SimpleStreamOptions["onPayload"];
290
291
  #onFinalPayload?: SimpleStreamOptions["onPayload"];
292
+ #checkpointTimestamps = new Set<number>();
293
+ #onAssistantCheckpoint?: (message: AssistantMessage) => Promise<void>;
291
294
  #onAssistantMessageEvent?: (message: AssistantMessage, event: AssistantMessageEvent) => void;
292
295
 
293
296
  /** Buffered Cursor tool results with text length at time of call (for correct ordering) */
@@ -327,6 +330,7 @@ export class Agent {
327
330
  this.#intentTracing = opts.intentTracing === true;
328
331
  this.#getToolChoice = opts.getToolChoice;
329
332
  this.#onAssistantMessageEvent = opts.onAssistantMessageEvent;
333
+ this.#onAssistantCheckpoint = opts.onAssistantCheckpoint;
330
334
  }
331
335
 
332
336
  /**
@@ -458,6 +462,10 @@ export class Agent {
458
462
  return () => this.#listeners.delete(fn);
459
463
  }
460
464
 
465
+ setAssistantCheckpointHandler(fn: (message: AssistantMessage) => Promise<void>): void {
466
+ this.#onAssistantCheckpoint = fn;
467
+ }
468
+
461
469
  setAssistantMessageEventInterceptor(
462
470
  fn: ((message: AssistantMessage, event: AssistantMessageEvent) => void) | undefined,
463
471
  ): void {
@@ -470,10 +478,29 @@ export class Agent {
470
478
  case "message_update":
471
479
  this.#state.streamMessage = event.message;
472
480
  break;
473
- case "message_end":
481
+ case "assistant_checkpoint": {
482
+ this.#checkpointTimestamps.add(event.message.timestamp);
483
+ const index = this.#state.messages.findIndex(
484
+ m => m.role === "assistant" && m.timestamp === event.message.timestamp,
485
+ );
486
+ if (index >= 0) this.#state.messages[index] = event.message;
487
+ else this.appendMessage(event.message);
488
+ break;
489
+ }
490
+ case "message_end": {
474
491
  this.#state.streamMessage = null;
475
- this.appendMessage(event.message);
492
+ const index =
493
+ event.message.role === "assistant" &&
494
+ (this.#checkpointTimestamps.has(event.message.timestamp) || !!event.message.interruption)
495
+ ? this.#state.messages.findIndex(
496
+ m => m.role === "assistant" && m.timestamp === event.message.timestamp,
497
+ )
498
+ : -1;
499
+ if (index >= 0) this.#state.messages[index] = event.message;
500
+ else this.appendMessage(event.message);
501
+ this.#checkpointTimestamps.delete(event.message.timestamp);
476
502
  break;
503
+ }
477
504
  case "tool_execution_start": {
478
505
  const pending = new Set(this.#state.pendingToolCalls);
479
506
  pending.add(event.toolCallId);
@@ -840,6 +867,7 @@ export class Agent {
840
867
  transformToolCallArguments: this.#transformToolCallArguments,
841
868
  intentTracing: this.#intentTracing,
842
869
  onAssistantMessageEvent: this.#onAssistantMessageEvent,
870
+ onAssistantCheckpoint: this.#onAssistantCheckpoint,
843
871
  getToolChoice,
844
872
  waitForSteeringMessages: signal => {
845
873
  if (this.#steeringQueue.length || signal.aborted) return Promise.resolve();
@@ -889,6 +917,15 @@ export class Agent {
889
917
  this.#state.streamMessage = event.message;
890
918
  break;
891
919
 
920
+ case "assistant_checkpoint": {
921
+ this.#checkpointTimestamps.add(event.message.timestamp);
922
+ const index = this.#state.messages.findIndex(
923
+ m => m.role === "assistant" && m.timestamp === event.message.timestamp,
924
+ );
925
+ if (index >= 0) this.#state.messages[index] = event.message;
926
+ else this.appendMessage(event.message);
927
+ break;
928
+ }
892
929
  case "message_end":
893
930
  partial = null;
894
931
  // Check if this is an assistant message with buffered Cursor tool results.
@@ -898,7 +935,27 @@ export class Agent {
898
935
  continue; // Skip default emit - split method handles everything
899
936
  }
900
937
  this.#state.streamMessage = null;
901
- this.appendMessage(event.message);
938
+ if (event.message.role === "assistant") {
939
+ const index =
940
+ this.#checkpointTimestamps.has(event.message.timestamp) || event.message.interruption
941
+ ? this.#state.messages.findIndex(
942
+ m => m.role === "assistant" && m.timestamp === event.message.timestamp,
943
+ )
944
+ : -1;
945
+ const message = event.message.interruption
946
+ ? {
947
+ ...event.message,
948
+ content: event.message.content.filter(
949
+ (_c, i) =>
950
+ event.message.role === "assistant" &&
951
+ event.message.interruption!.completedContentIndices.includes(i),
952
+ ),
953
+ }
954
+ : event.message;
955
+ if (index >= 0) this.#state.messages[index] = message;
956
+ else this.appendMessage(message);
957
+ } else this.appendMessage(event.message);
958
+ this.#checkpointTimestamps.delete(event.message.timestamp);
902
959
  break;
903
960
 
904
961
  case "tool_execution_start": {
package/src/proxy.ts CHANGED
@@ -43,17 +43,26 @@ export type ProxyAssistantMessageEvent =
43
43
  | { type: "thinking_end"; contentIndex: number; contentSignature?: string }
44
44
  | { type: "toolcall_start"; contentIndex: number; id: string; toolName: string }
45
45
  | { type: "toolcall_delta"; contentIndex: number; delta: string }
46
- | { type: "toolcall_end"; contentIndex: number }
46
+ | {
47
+ type: "toolcall_end";
48
+ contentIndex: number;
49
+ toolCall?: ToolCall;
50
+ providerPayload?: AssistantMessage["providerPayload"];
51
+ }
47
52
  | {
48
53
  type: "done";
49
54
  reason: Extract<StopReason, "stop" | "length" | "toolUse">;
50
55
  usage: AssistantMessage["usage"];
56
+ providerPayload?: AssistantMessage["providerPayload"];
57
+ interruption?: AssistantMessage["interruption"];
51
58
  }
52
59
  | {
53
60
  type: "error";
54
61
  reason: Extract<StopReason, "aborted" | "error">;
55
62
  errorMessage?: string;
56
63
  usage: AssistantMessage["usage"];
64
+ providerPayload?: AssistantMessage["providerPayload"];
65
+ interruption?: AssistantMessage["interruption"];
57
66
  };
58
67
 
59
68
  export interface ProxyStreamOptions extends SimpleStreamOptions {
@@ -135,6 +144,7 @@ export function streamProxy(model: Model, context: Context, options: ProxyStream
135
144
  repetitionPenalty: options.repetitionPenalty,
136
145
  maxTokens: options.maxTokens,
137
146
  reasoning: options.reasoning,
147
+ maxRetries: options.maxRetries,
138
148
  },
139
149
  }),
140
150
  signal: options.signal,
@@ -307,6 +317,8 @@ function processProxyEvent(
307
317
  }
308
318
 
309
319
  case "toolcall_end": {
320
+ if (proxyEvent.providerPayload) partial.providerPayload = proxyEvent.providerPayload;
321
+ if (proxyEvent.toolCall) partial.content[proxyEvent.contentIndex] = proxyEvent.toolCall;
310
322
  const content = partial.content[proxyEvent.contentIndex];
311
323
  if (content?.type === "toolCall") {
312
324
  delete (content as any).partialJson;
@@ -323,12 +335,16 @@ function processProxyEvent(
323
335
  case "done":
324
336
  partial.stopReason = proxyEvent.reason;
325
337
  partial.usage = proxyEvent.usage;
338
+ partial.providerPayload = proxyEvent.providerPayload ?? partial.providerPayload;
339
+ partial.interruption = proxyEvent.interruption;
326
340
  return { type: "done", reason: proxyEvent.reason, message: partial };
327
341
 
328
342
  case "error":
329
343
  partial.stopReason = proxyEvent.reason;
330
344
  partial.errorMessage = proxyEvent.errorMessage;
331
345
  partial.usage = proxyEvent.usage;
346
+ partial.providerPayload = proxyEvent.providerPayload ?? partial.providerPayload;
347
+ partial.interruption = proxyEvent.interruption;
332
348
  return { type: "error", reason: proxyEvent.reason, error: partial };
333
349
  }
334
350
  }
package/src/types.ts CHANGED
@@ -143,6 +143,7 @@ export interface AgentLoopConfig extends SimpleStreamOptions {
143
143
  * Inspect assistant streaming events before they are published to the outer agent event stream.
144
144
  * Callers may abort synchronously to stop consuming buffered provider events.
145
145
  */
146
+ onAssistantCheckpoint?: (message: AssistantMessage) => Promise<void>;
146
147
  onAssistantMessageEvent?: (message: AssistantMessage, event: AssistantMessageEvent) => void;
147
148
 
148
149
  /**
@@ -299,12 +300,13 @@ export type AgentEvent =
299
300
  | { type: "agent_end"; messages: AgentMessage[] }
300
301
  // Turn lifecycle - a turn is one assistant response + any tool calls/results
301
302
  | { type: "turn_start" }
302
- | { type: "turn_end"; message: AgentMessage; toolResults: ToolResultMessage[] }
303
+ | { type: "turn_end"; message: AgentMessage; toolResults: ToolResultMessage[]; toolChoiceServed?: boolean }
303
304
  // Message lifecycle - emitted for user, assistant, and toolResult messages
304
305
  | { type: "message_start"; message: AgentMessage }
305
306
  // Only emitted for assistant messages during streaming
306
307
  | { type: "message_update"; message: AgentMessage; assistantMessageEvent: AssistantMessageEvent }
307
308
  | { type: "message_end"; message: AgentMessage }
309
+ | { type: "assistant_checkpoint"; message: AssistantMessage }
308
310
  // Tool execution lifecycle
309
311
  | {
310
312
  type: "tool_execution_start";