@tanstack/ai-persistence 0.1.5 → 0.4.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.
package/src/middleware.ts CHANGED
@@ -1,15 +1,29 @@
1
1
  import {
2
2
  defineChatMiddleware,
3
3
  getDetachableRun,
4
+ InterruptResumeValidationError,
5
+ readInterruptBinding,
6
+ validateInterruptResumeBatch,
4
7
  wasCancelRequested,
5
8
  } from '@tanstack/ai'
6
- import { providePendingTurn } from '@tanstack/ai/adapter-internals'
9
+ import {
10
+ createInterruptBinding,
11
+ getGenericInterruptDefinitionRegistry,
12
+ providePendingTurn,
13
+ rehydrateInterruptRequest,
14
+ } from '@tanstack/ai/adapter-internals'
15
+ import type {
16
+ GenericInterruptRequest,
17
+ InterruptDefinition,
18
+ } from '@tanstack/ai/adapter-internals'
7
19
  import { base64ToUint8Array } from '@tanstack/ai-utils'
8
20
  import {
9
21
  InterruptsCapability,
10
22
  PersistenceCapability,
23
+ PersistenceCompletionCapability,
11
24
  provideInterrupts,
12
25
  providePersistence,
26
+ providePersistenceCompletion,
13
27
  } from './capabilities'
14
28
  import {
15
29
  validateChatPersistenceStores,
@@ -28,13 +42,16 @@ import type {
28
42
  GenerationFinishInfo,
29
43
  GenerationMiddleware,
30
44
  GenerationMiddlewareContext,
31
- ModelMessage,
45
+ Interrupt,
46
+ PendingInterruptResumeRecord,
32
47
  PersistedArtifactActivity,
33
48
  PersistedArtifactRef,
34
49
  PersistedArtifactRole,
35
50
  RunAgentResumeItem,
36
51
  StreamChunk,
52
+ Tool,
37
53
  ToolApprovalResolution,
54
+ BilledUsage,
38
55
  TokenUsage,
39
56
  } from '@tanstack/ai'
40
57
  import type {
@@ -43,6 +60,7 @@ import type {
43
60
  ArtifactRecord,
44
61
  BlobBody,
45
62
  ChatTranscriptStores,
63
+ InterruptCommitEntry,
46
64
  InterruptRecord,
47
65
  RunStore,
48
66
  } from './types'
@@ -268,33 +286,153 @@ interface RunStateEntry {
268
286
  */
269
287
  streamingMessageId?: string
270
288
  streamingMessageCreatedAt?: Date
289
+ completion?: {
290
+ promise: Promise<void>
291
+ resolve: () => void
292
+ reject: (error: unknown) => void
293
+ }
271
294
  }
272
295
 
273
296
  const runState = new WeakMap<object, RunStateEntry>()
274
297
 
275
298
  const validResumeStatuses = new Set(['resolved', 'cancelled'])
276
299
 
300
+ function mergeMaps<K, V>(
301
+ left?: ReadonlyMap<K, V>,
302
+ right?: ReadonlyMap<K, V>,
303
+ ): Map<K, V> | undefined {
304
+ if (!left && !right) return undefined
305
+ return new Map([...(left ?? []), ...(right ?? [])])
306
+ }
307
+
308
+ function mergeSets<T>(
309
+ left?: ReadonlySet<T>,
310
+ right?: ReadonlySet<T>,
311
+ ): Set<T> | undefined {
312
+ if (!left && !right) return undefined
313
+ return new Set([...(left ?? []), ...(right ?? [])])
314
+ }
315
+
316
+ function mergeResumeToolState(
317
+ left: ChatResumeToolState | undefined,
318
+ right: ChatResumeToolState | undefined,
319
+ ): ChatResumeToolState | undefined {
320
+ if (!left) return right
321
+ if (!right) return left
322
+ return {
323
+ approvals: mergeMaps(left.approvals, right.approvals),
324
+ clientToolResults: mergeMaps(
325
+ left.clientToolResults,
326
+ right.clientToolResults,
327
+ ),
328
+ genericInterrupts: mergeMaps(
329
+ left.genericInterrupts,
330
+ right.genericInterrupts,
331
+ ),
332
+ genericInterruptRequests: mergeMaps(
333
+ left.genericInterruptRequests,
334
+ right.genericInterruptRequests,
335
+ ),
336
+ deniedToolResults: mergeMaps(
337
+ left.deniedToolResults,
338
+ right.deniedToolResults,
339
+ ),
340
+ cancelledToolCallIds: mergeSets(
341
+ left.cancelledToolCallIds,
342
+ right.cancelledToolCallIds,
343
+ ),
344
+ }
345
+ }
346
+
347
+ function rejectMixedRunPending(
348
+ pending: Array<InterruptRecord>,
349
+ ctx: Pick<ChatMiddlewareContext, 'threadId' | 'runId'>,
350
+ ): void {
351
+ const runIds = new Set(pending.map((interrupt) => interrupt.runId))
352
+ if (runIds.size <= 1) return
353
+ throw new InterruptResumeValidationError([
354
+ {
355
+ scope: 'batch',
356
+ threadId: ctx.threadId,
357
+ interruptedRunId: ctx.runId,
358
+ generation: 0,
359
+ interruptIds: pending.map((interrupt) => interrupt.interruptId),
360
+ code: 'stale',
361
+ message: 'Thread has pending interrupts from more than one run.',
362
+ source: 'server',
363
+ retryable: false,
364
+ },
365
+ ])
366
+ }
367
+
277
368
  function validatePendingResumes(
278
369
  pending: Array<InterruptRecord>,
279
370
  resume: Array<RunAgentResumeItem> | undefined,
371
+ ctx: Pick<ChatMiddlewareContext, 'threadId' | 'runId'>,
280
372
  ): Map<string, RunAgentResumeItem> {
373
+ const interruptedRunId = pending[0]?.runId ?? ctx.runId
374
+ const failure = (
375
+ interruptId: string,
376
+ code: 'conflict' | 'unknown-interrupt',
377
+ message: string,
378
+ ): never => {
379
+ throw new InterruptResumeValidationError([
380
+ {
381
+ scope: 'item',
382
+ threadId: ctx.threadId,
383
+ interruptedRunId,
384
+ generation: 0,
385
+ interruptId,
386
+ code,
387
+ message,
388
+ source: 'client',
389
+ retryable: false,
390
+ },
391
+ {
392
+ scope: 'batch',
393
+ threadId: ctx.threadId,
394
+ interruptedRunId,
395
+ generation: 0,
396
+ interruptIds: pending.map((interrupt) => interrupt.interruptId),
397
+ code: code === 'conflict' ? 'conflict' : 'incomplete-batch',
398
+ message:
399
+ 'Resume entries must resolve or cancel the complete interrupt batch.',
400
+ source: 'client',
401
+ retryable: false,
402
+ },
403
+ ])
404
+ }
281
405
  const pendingInterruptIds = new Set(
282
406
  pending.map((interrupt) => interrupt.interruptId),
283
407
  )
284
- const resumeByInterruptId = new Map(
285
- (resume ?? []).map((entry) => [entry.interruptId, entry]),
286
- )
408
+ const resumeByInterruptId = new Map<string, RunAgentResumeItem>()
409
+ for (const entry of resume ?? []) {
410
+ if (resumeByInterruptId.has(entry.interruptId)) {
411
+ return failure(
412
+ entry.interruptId,
413
+ 'conflict',
414
+ `Interrupt ${entry.interruptId} has duplicate resume entries.`,
415
+ )
416
+ }
417
+ resumeByInterruptId.set(entry.interruptId, entry)
418
+ }
287
419
  if (pending.length === 0) {
288
420
  const staleEntry = resume?.[0]
289
421
  if (staleEntry) {
290
- throw new Error(
422
+ return failure(
423
+ staleEntry.interruptId,
424
+ 'unknown-interrupt',
291
425
  `Resume entry references non-pending interrupt ${staleEntry.interruptId}.`,
292
426
  )
293
427
  }
294
428
  return resumeByInterruptId
295
429
  }
430
+ const firstPending = pending[0]
431
+ if (firstPending === undefined) return resumeByInterruptId
296
432
  if (!resume || resume.length === 0) {
297
- throw new Error(
433
+ return failure(
434
+ firstPending.interruptId,
435
+ 'unknown-interrupt',
298
436
  `Thread has pending interrupts; resume is required before accepting new input.`,
299
437
  )
300
438
  }
@@ -302,19 +440,25 @@ function validatePendingResumes(
302
440
  for (const interrupt of pending) {
303
441
  const entry = resumeByInterruptId.get(interrupt.interruptId)
304
442
  if (!entry) {
305
- throw new Error(
443
+ return failure(
444
+ interrupt.interruptId,
445
+ 'unknown-interrupt',
306
446
  `Missing resume entry for pending interrupt ${interrupt.interruptId}.`,
307
447
  )
308
448
  }
309
449
  if (!validResumeStatuses.has(entry.status)) {
310
- throw new Error(
450
+ return failure(
451
+ interrupt.interruptId,
452
+ 'unknown-interrupt',
311
453
  `Invalid resume status for pending interrupt ${interrupt.interruptId}: ${entry.status}.`,
312
454
  )
313
455
  }
314
456
  }
315
457
  for (const entry of resume) {
316
458
  if (!pendingInterruptIds.has(entry.interruptId)) {
317
- throw new Error(
459
+ return failure(
460
+ entry.interruptId,
461
+ 'unknown-interrupt',
318
462
  `Resume entry references non-pending interrupt ${entry.interruptId}.`,
319
463
  )
320
464
  }
@@ -327,13 +471,52 @@ async function applyPendingResumes(
327
471
  resumeByInterruptId: Map<string, RunAgentResumeItem>,
328
472
  interrupts: NonNullable<AIPersistence['stores']['interrupts']>,
329
473
  ): Promise<void> {
474
+ const entries: Array<InterruptCommitEntry> = []
330
475
  for (const interrupt of pending) {
331
476
  const entry = resumeByInterruptId.get(interrupt.interruptId)
332
477
  if (!entry) continue
333
478
  if (entry.status === 'resolved') {
334
- await interrupts.resolve(interrupt.interruptId, entry.payload)
479
+ entries.push({
480
+ interruptId: interrupt.interruptId,
481
+ status: 'resolved',
482
+ response: entry.payload,
483
+ })
335
484
  } else {
336
- await interrupts.cancel(interrupt.interruptId)
485
+ entries.push({
486
+ interruptId: interrupt.interruptId,
487
+ status: 'cancelled',
488
+ })
489
+ }
490
+ }
491
+ if (interrupts.commitBatch) {
492
+ await interrupts.commitBatch(entries)
493
+ return
494
+ }
495
+ const ids = new Set<string>()
496
+ for (const entry of entries) {
497
+ if (ids.has(entry.interruptId)) {
498
+ throw new Error(
499
+ `Interrupt batch contains duplicate id: ${entry.interruptId}.`,
500
+ )
501
+ }
502
+ ids.add(entry.interruptId)
503
+ const existing = await interrupts.get(entry.interruptId)
504
+ if (!existing) {
505
+ throw new Error(
506
+ `Interrupt batch references missing id: ${entry.interruptId}.`,
507
+ )
508
+ }
509
+ if (existing.status !== 'pending') {
510
+ throw new Error(
511
+ `Interrupt batch references non-pending id: ${entry.interruptId}.`,
512
+ )
513
+ }
514
+ }
515
+ for (const entry of entries) {
516
+ if (entry.status === 'resolved') {
517
+ await interrupts.resolve(entry.interruptId, entry.response)
518
+ } else {
519
+ await interrupts.cancel(entry.interruptId)
337
520
  }
338
521
  }
339
522
  }
@@ -377,6 +560,267 @@ function interruptKind(interrupt: InterruptRecord): string | undefined {
377
560
  return metadata ? stringField(metadata, 'kind') : undefined
378
561
  }
379
562
 
563
+ function hasReservedInterruptBinding(payload: unknown): boolean {
564
+ const descriptor = objectValue(payload)
565
+ const metadata = objectValue(descriptor?.metadata)
566
+ return !!metadata && 'tanstack:interruptBinding' in metadata
567
+ }
568
+
569
+ function isPersistedInterruptDescriptor(
570
+ value: unknown,
571
+ ): value is Interrupt & { reason: string; message: string } {
572
+ const record = objectValue(value)
573
+ return (
574
+ !!record &&
575
+ typeof record.id === 'string' &&
576
+ typeof record.reason === 'string' &&
577
+ typeof record.message === 'string'
578
+ )
579
+ }
580
+
581
+ /**
582
+ * Does this pending record belong to the TanStack chat resume protocol?
583
+ *
584
+ * An external system can persist an AG-UI descriptor in the same durable
585
+ * thread. A descriptor without a TanStack binding or legacy tool marker stays
586
+ * pending for its owner, but it does not make this resume incomplete. Older
587
+ * opaque records remain owned because their provenance cannot be known.
588
+ */
589
+ function isChatOwnedPendingInterrupt(interrupt: InterruptRecord): boolean {
590
+ const kind = interruptKind(interrupt)
591
+ return (
592
+ !isPersistedInterruptDescriptor(interrupt.payload) ||
593
+ stringField(interrupt.payload, 'toolCallId') !== undefined ||
594
+ kind === 'approval' ||
595
+ kind === 'client_tool' ||
596
+ hasReservedInterruptBinding(interrupt.payload)
597
+ )
598
+ }
599
+
600
+ function durableGenericFailure(
601
+ ctx: Pick<ChatMiddlewareContext, 'threadId' | 'runId'>,
602
+ persisted: InterruptRecord,
603
+ message: string,
604
+ ): InterruptResumeValidationError {
605
+ return new InterruptResumeValidationError([
606
+ {
607
+ scope: 'item',
608
+ threadId: ctx.threadId,
609
+ interruptedRunId: persisted.runId || ctx.runId,
610
+ generation: 0,
611
+ interruptId: persisted.interruptId,
612
+ code: 'stale',
613
+ message,
614
+ source: 'server',
615
+ retryable: false,
616
+ },
617
+ {
618
+ scope: 'batch',
619
+ threadId: ctx.threadId,
620
+ interruptedRunId: persisted.runId || ctx.runId,
621
+ generation: 0,
622
+ interruptIds: [persisted.interruptId],
623
+ code: 'item-validation-failed',
624
+ message: 'One or more persisted interrupt records are invalid.',
625
+ source: 'server',
626
+ retryable: false,
627
+ },
628
+ ])
629
+ }
630
+
631
+ async function durableGenericResumeState(
632
+ ctx: ChatMiddlewareContext,
633
+ pending: Array<InterruptRecord>,
634
+ resume: ReadonlyArray<RunAgentResumeItem>,
635
+ tools: Array<Tool>,
636
+ ): Promise<ChatResumeToolState | undefined> {
637
+ const registry = getGenericInterruptDefinitionRegistry(ctx, {
638
+ optional: true,
639
+ })
640
+ const records: Array<PendingInterruptResumeRecord> = []
641
+
642
+ for (const persisted of pending) {
643
+ if (!isPersistedInterruptDescriptor(persisted.payload)) {
644
+ if (hasReservedInterruptBinding(persisted.payload)) {
645
+ throw durableGenericFailure(
646
+ ctx,
647
+ persisted,
648
+ `Persisted interrupt ${persisted.interruptId} has an invalid binding descriptor.`,
649
+ )
650
+ }
651
+ continue
652
+ }
653
+ const descriptor = persisted.payload
654
+ const binding = readInterruptBinding(descriptor)
655
+ if (!binding) {
656
+ if (hasReservedInterruptBinding(descriptor)) {
657
+ throw durableGenericFailure(
658
+ ctx,
659
+ persisted,
660
+ `Persisted interrupt ${persisted.interruptId} has an invalid or incomplete binding.`,
661
+ )
662
+ }
663
+ continue
664
+ }
665
+ if (
666
+ descriptor.id !== persisted.interruptId ||
667
+ binding.interruptId !== persisted.interruptId ||
668
+ binding.interruptedRunId !== persisted.runId ||
669
+ binding.generation !== 0
670
+ ) {
671
+ throw durableGenericFailure(
672
+ ctx,
673
+ persisted,
674
+ `Persisted interrupt ${persisted.interruptId} has stale correlation metadata.`,
675
+ )
676
+ }
677
+ if (binding.kind !== 'generic') {
678
+ records.push({
679
+ interruptId: persisted.interruptId,
680
+ payload: descriptor,
681
+ binding,
682
+ })
683
+ continue
684
+ }
685
+ if (
686
+ !binding.definitionId ||
687
+ !binding.key ||
688
+ binding.batchIndex === undefined
689
+ ) {
690
+ records.push({
691
+ interruptId: persisted.interruptId,
692
+ payload: descriptor,
693
+ binding,
694
+ })
695
+ continue
696
+ }
697
+ if (!registry) {
698
+ throw durableGenericFailure(
699
+ ctx,
700
+ persisted,
701
+ `Persisted generic interrupt ${persisted.interruptId} cannot be restored because no interrupt registry is available.`,
702
+ )
703
+ }
704
+ const definition = registry.definitions.get(binding.definitionId)
705
+ if (!definition) {
706
+ throw durableGenericFailure(
707
+ ctx,
708
+ persisted,
709
+ `Persisted generic interrupt definition ${binding.definitionId} is unavailable.`,
710
+ )
711
+ }
712
+ const metadata = objectValue(descriptor.metadata)
713
+ const payload = metadata?.['tanstack:interruptPayload']
714
+ let request: GenericInterruptRequest<
715
+ InterruptDefinition<any, any, any, any>
716
+ >
717
+ try {
718
+ request = rehydrateInterruptRequest(definition, {
719
+ key: binding.key,
720
+ reason: descriptor.reason,
721
+ message: descriptor.message,
722
+ ...(descriptor.expiresAt !== undefined
723
+ ? { expiresAt: descriptor.expiresAt }
724
+ : {}),
725
+ ...(payload !== undefined ? { payload } : {}),
726
+ })
727
+ } catch (error) {
728
+ throw durableGenericFailure(
729
+ ctx,
730
+ persisted,
731
+ `Persisted generic interrupt ${persisted.interruptId} is invalid: ${error instanceof Error ? error.message : String(error)}`,
732
+ )
733
+ }
734
+ const emitted = createInterruptBinding(request, {
735
+ batchIndex: binding.batchIndex,
736
+ })
737
+ if (
738
+ emitted.descriptor.responseSchemaHash !== binding.responseSchemaHash ||
739
+ emitted.descriptor.payloadSchemaHash !== binding.payloadSchemaHash ||
740
+ binding.interruptId !== persisted.interruptId
741
+ ) {
742
+ throw durableGenericFailure(
743
+ ctx,
744
+ persisted,
745
+ `Persisted generic interrupt ${persisted.interruptId} is stale.`,
746
+ )
747
+ }
748
+ records.push({
749
+ interruptId: persisted.interruptId,
750
+ payload: descriptor,
751
+ binding,
752
+ genericRequest: request,
753
+ })
754
+ }
755
+
756
+ const firstRecord = records[0]
757
+ if (firstRecord === undefined) return undefined
758
+ const interruptedRunId = firstRecord.binding.interruptedRunId
759
+ const generation = firstRecord.binding.generation
760
+ const validated = await validateInterruptResumeBatch({
761
+ threadId: ctx.threadId,
762
+ interruptedRunId,
763
+ generation,
764
+ pending: records,
765
+ resume: resume.filter((entry) =>
766
+ records.some((record) => record.interruptId === entry.interruptId),
767
+ ),
768
+ tools,
769
+ })
770
+ if (validated.errors.length > 0 || !validated.resumeToolState) {
771
+ throw new InterruptResumeValidationError(validated.errors)
772
+ }
773
+ type GenericRecord = PendingInterruptResumeRecord & {
774
+ binding: Extract<
775
+ PendingInterruptResumeRecord['binding'],
776
+ { kind: 'generic' }
777
+ >
778
+ genericRequest: GenericInterruptRequest<
779
+ InterruptDefinition<any, any, any, any>
780
+ >
781
+ }
782
+ const isGenericRecord = (
783
+ record: PendingInterruptResumeRecord,
784
+ ): record is GenericRecord =>
785
+ record.binding.kind === 'generic' && record.genericRequest !== undefined
786
+ const genericRecords: Array<{ record: GenericRecord; batchIndex: number }> =
787
+ []
788
+ const batchIndexes = new Set<number>()
789
+ for (const record of records) {
790
+ if (!isGenericRecord(record)) continue
791
+ const batchIndex = record.binding.batchIndex
792
+ if (batchIndex === undefined || batchIndexes.has(batchIndex)) {
793
+ throw new InterruptResumeValidationError([
794
+ {
795
+ scope: 'batch',
796
+ threadId: ctx.threadId,
797
+ interruptedRunId,
798
+ generation,
799
+ interruptIds: records.map((item) => item.interruptId),
800
+ code: 'stale',
801
+ message:
802
+ 'Persisted generic interrupts have duplicate or invalid batch indexes.',
803
+ source: 'server',
804
+ retryable: false,
805
+ },
806
+ ])
807
+ }
808
+ batchIndexes.add(batchIndex)
809
+ genericRecords.push({ record, batchIndex })
810
+ }
811
+ genericRecords.sort((left, right) => left.batchIndex - right.batchIndex)
812
+ return {
813
+ ...validated.resumeToolState,
814
+ genericInterruptRequests: new Map(
815
+ genericRecords.flatMap(({ record }) =>
816
+ record.genericRequest
817
+ ? [[record.interruptId, record.genericRequest] as const]
818
+ : [],
819
+ ),
820
+ ),
821
+ }
822
+ }
823
+
380
824
  function resolvedApprovalDecision(entry: RunAgentResumeItem): boolean {
381
825
  if (entry.status === 'cancelled') return false
382
826
  const payload = objectValue(entry.payload)
@@ -437,43 +881,6 @@ function resumeToolStateFromPending(
437
881
  return { approvals, clientToolResults, cancelledToolCallIds }
438
882
  }
439
883
 
440
- /**
441
- * Build the transcript to persist when a run finishes successfully.
442
- *
443
- * The chat engine appends an assistant message to the middleware message list
444
- * only when that turn carries tool calls (to feed the agent loop); a run's
445
- * terminal *text* reply is never appended. So `ctx.messages` at `onFinish` is
446
- * missing the assistant's final answer. Reattach it from the finish info —
447
- * `info.content` is the last turn's accumulated text (reset each cycle) — so
448
- * the stored thread is the complete conversation a server-authoritative client
449
- * hydrates on load. A guard avoids duplicating a terminal assistant turn should
450
- * the engine ever start appending it itself.
451
- */
452
- function finishedTranscript(
453
- messages: ReadonlyArray<ModelMessage>,
454
- info: FinishInfo,
455
- messageId: string | undefined,
456
- createdAt: Date | undefined,
457
- ): Array<ModelMessage> {
458
- const transcript = [...messages]
459
- const last = transcript[transcript.length - 1]
460
- const alreadyPresent =
461
- last?.role === 'assistant' &&
462
- last.toolCalls === undefined &&
463
- last.content === info.content
464
- if (info.content && !alreadyPresent) {
465
- // Stamp the terminal turn with its stream messageId so a hydrated bubble
466
- // keeps the same identity as the live stream (in-place resume on reload).
467
- transcript.push({
468
- role: 'assistant',
469
- content: info.content,
470
- ...(messageId ? { id: messageId } : {}),
471
- ...(createdAt ? { createdAt } : {}),
472
- })
473
- }
474
- return transcript
475
- }
476
-
477
884
  function interruptPayload(interrupt: unknown): Record<string, unknown> {
478
885
  return interrupt && typeof interrupt === 'object'
479
886
  ? { ...(interrupt as Record<string, unknown>) }
@@ -1349,6 +1756,7 @@ function accumulateTokenUsage(
1349
1756
  next.durationSeconds,
1350
1757
  )
1351
1758
  const unitsBilled = sumOptionalNumber(current.unitsBilled, next.unitsBilled)
1759
+ const billed = accumulateBilled(current.billed, next.billed)
1352
1760
  const cost = sumOptionalNumber(current.cost, next.cost)
1353
1761
 
1354
1762
  return {
@@ -1361,12 +1769,27 @@ function accumulateTokenUsage(
1361
1769
  ...(completionTokensDetails ? { completionTokensDetails } : {}),
1362
1770
  ...(durationSeconds !== undefined ? { durationSeconds } : {}),
1363
1771
  ...(unitsBilled !== undefined ? { unitsBilled } : {}),
1772
+ ...(billed !== undefined ? { billed } : {}),
1364
1773
  ...(cost !== undefined ? { cost } : {}),
1365
1774
  ...(costDetails ? { costDetails } : {}),
1366
1775
  ...(providerUsageDetails ? { providerUsageDetails } : {}),
1367
1776
  }
1368
1777
  }
1369
1778
 
1779
+ /**
1780
+ * Sum billed quantities when both reports use the same unit. Different units
1781
+ * cannot be added, so the later report wins.
1782
+ */
1783
+ function accumulateBilled(
1784
+ current: BilledUsage | undefined,
1785
+ next: BilledUsage | undefined,
1786
+ ): BilledUsage | undefined {
1787
+ if (!current) return next
1788
+ if (!next) return current
1789
+ if (current.unit !== next.unit) return next
1790
+ return { quantity: current.quantity + next.quantity, unit: current.unit }
1791
+ }
1792
+
1370
1793
  async function completeRun(
1371
1794
  runs: RunStore | undefined,
1372
1795
  runId: string,
@@ -1510,6 +1933,7 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1510
1933
 
1511
1934
  const provides = [
1512
1935
  PersistenceCapability,
1936
+ PersistenceCompletionCapability,
1513
1937
  ...(wantsInterrupts ? [InterruptsCapability] : []),
1514
1938
  ]
1515
1939
 
@@ -1519,9 +1943,27 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1519
1943
  setup(ctx: ChatMiddlewareContext) {
1520
1944
  providePersistence(ctx, persistence)
1521
1945
 
1946
+ let resolveCompletion: () => void = () => undefined
1947
+ let rejectCompletion: (error: unknown) => void = () => undefined
1948
+ const completion = new Promise<void>((resolve, reject) => {
1949
+ resolveCompletion = resolve
1950
+ rejectCompletion = reject
1951
+ })
1952
+ // Consumers may not need this capability. Mark the rejection handled while
1953
+ // preserving the original promise for callers that do await it.
1954
+ void completion.catch(() => undefined)
1955
+
1522
1956
  runState.set(ctx, {
1523
1957
  merged: false,
1524
1958
  interrupted: false,
1959
+ completion: {
1960
+ promise: completion,
1961
+ resolve: resolveCompletion,
1962
+ reject: rejectCompletion,
1963
+ },
1964
+ })
1965
+ providePersistenceCompletion(ctx, {
1966
+ waitForRunCompletion: () => completion,
1525
1967
  })
1526
1968
 
1527
1969
  if (wantsInterrupts && persistence.stores.interrupts) {
@@ -1558,11 +2000,15 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1558
2000
  const pending = await persistence.stores.interrupts.listPending(
1559
2001
  ctx.threadId,
1560
2002
  )
1561
- // Gate: a thread with pending interrupts must carry a resume batch that
1562
- // references them.
2003
+ // Gate only records that this chat owns. A foreign AG-UI interrupt can
2004
+ // share the durable thread, but its owner resolves it outside this
2005
+ // resume protocol. Including it would deadlock this chat resume.
2006
+ const ownedPending = pending.filter(isChatOwnedPendingInterrupt)
2007
+ rejectMixedRunPending(ownedPending, ctx)
1563
2008
  const resumeByInterruptId = validatePendingResumes(
1564
- pending,
2009
+ ownedPending,
1565
2010
  config.resume,
2011
+ ctx,
1566
2012
  )
1567
2013
  // Persistence is the server-authoritative resume path: translate the
1568
2014
  // persisted interrupts into the engine's resume tool state and CLEAR
@@ -1571,18 +2017,29 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1571
2017
  // persistence flow deliberately omits).
1572
2018
  if ((config.resume?.length ?? 0) > 0) {
1573
2019
  const resumeToolState = resumeToolStateFromPending(
1574
- pending,
2020
+ ownedPending,
1575
2021
  resumeByInterruptId,
1576
2022
  )
2023
+ const genericResumeState = await durableGenericResumeState(
2024
+ ctx,
2025
+ ownedPending,
2026
+ config.resume ?? [],
2027
+ config.tools,
2028
+ )
1577
2029
  patch.resume = []
1578
- if (resumeToolState) patch.resumeToolState = resumeToolState
2030
+ if (resumeToolState || genericResumeState) {
2031
+ patch.resumeToolState = mergeResumeToolState(
2032
+ resumeToolState,
2033
+ genericResumeState,
2034
+ )
2035
+ }
1579
2036
  }
1580
2037
  // Defer marking these interrupts resolved/cancelled until the run
1581
2038
  // succeeds (see commitPendingResumes). Committing here would consume the
1582
2039
  // approval even if the run then failed, breaking a retry.
1583
2040
  const state = runState.get(ctx)
1584
- if (state && pending.length > 0) {
1585
- state.pendingResumes = { pending, resumeByInterruptId }
2041
+ if (state && ownedPending.length > 0) {
2042
+ state.pendingResumes = { pending: ownedPending, resumeByInterruptId }
1586
2043
  }
1587
2044
  }
1588
2045
 
@@ -1613,11 +2070,9 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1613
2070
  },
1614
2071
 
1615
2072
  async onChunk(ctx: ChatMiddlewareContext, chunk: StreamChunk) {
1616
- // Always capture the current assistant turn's stream messageId (cheap),
1617
- // regardless of snapshotStreaming — it's persisted onto the assistant
1618
- // message so its identity survives hydrate and a reload resumes the same
1619
- // bubble in place.
1620
- if (ctx.phase === 'modelStream') {
2073
+ // Capture the current assistant turn's identity for optional in-progress
2074
+ // snapshots. Completed messages already live in `ctx.messages`.
2075
+ if (snapshotStreaming && ctx.phase === 'modelStream') {
1621
2076
  const s = runState.get(ctx)
1622
2077
  if (s && chunk.type === 'TEXT_MESSAGE_START') {
1623
2078
  // An empty/malformed messageId means "no identity" (matching the
@@ -1644,9 +2099,8 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1644
2099
 
1645
2100
  // (B) Optional throttled snapshot of the in-progress assistant reply, so
1646
2101
  // partial output survives a crash/reload before onFinish. Off unless
1647
- // `snapshotStreaming` is set. We accumulate the terminal turn's text here
1648
- // (the engine only appends assistant turns with tool calls to
1649
- // `ctx.messages`, never a streaming text reply), then persist
2102
+ // `snapshotStreaming` is set. The completed turn enters `ctx.messages`
2103
+ // only after streaming ends, so accumulate its text here and persist
1650
2104
  // `ctx.messages` + that partial assistant message (tagged with its id).
1651
2105
  if (
1652
2106
  snapshotStreaming &&
@@ -1732,21 +2186,30 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1732
2186
  // resumes stay pending so a retry can re-apply them. Completing the run
1733
2187
  // or consuming approvals before the durable history lands leaves a
1734
2188
  // "finished" run whose transcript is missing the terminal turn.
1735
- await messageStore.saveThread(
1736
- ctx.threadId,
1737
- finishedTranscript(
1738
- ctx.messages,
1739
- info,
1740
- state?.streamingMessageId,
1741
- state?.streamingMessageCreatedAt,
1742
- ),
1743
- )
1744
- await completeRun(runs, ctx.runId, state?.usage ?? info.usage)
1745
- await commitPendingResumes(state, persistence.stores.interrupts)
2189
+ try {
2190
+ await messageStore.saveThread(ctx.threadId, [...ctx.messages])
2191
+ await commitPendingResumes(state, persistence.stores.interrupts)
2192
+ await completeRun(runs, ctx.runId, state?.usage ?? info.usage)
2193
+ state?.completion?.resolve()
2194
+ } catch (error) {
2195
+ // Core has already selected its terminal hook. Persist the failed run
2196
+ // here, so a failed transcript save or batch write does not leave an
2197
+ // interrupted or completed run whose pending records need retrying.
2198
+ try {
2199
+ await failRun(runs, ctx.runId, error, state?.usage)
2200
+ } finally {
2201
+ state?.completion?.reject(error)
2202
+ }
2203
+ throw error
2204
+ }
1746
2205
  },
1747
2206
 
1748
2207
  async onError(ctx: ChatMiddlewareContext, info: ErrorInfo) {
1749
- await failRun(runs, ctx.runId, info.error, runState.get(ctx)?.usage)
2208
+ try {
2209
+ await failRun(runs, ctx.runId, info.error, runState.get(ctx)?.usage)
2210
+ } finally {
2211
+ runState.get(ctx)?.completion?.reject(info.error)
2212
+ }
1750
2213
  },
1751
2214
 
1752
2215
  async onAbort(ctx: ChatMiddlewareContext, info: AbortInfo) {
@@ -1756,10 +2219,6 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1756
2219
  // (`info.cancelRequested`, set when the cancel aborted this host's signal)
1757
2220
  // and durable (`RunRecord.cancelRequested`, the only channel that reaches
1758
2221
  // a run being driven elsewhere).
1759
- const cancelled =
1760
- info.cancelRequested === true ||
1761
- (runs !== undefined && (await wasCancelRequested(runs, ctx.runId)))
1762
-
1763
2222
  // A run paused at an interrupt boundary is waiting for a HUMAN, not for
1764
2223
  // this socket. `chat()` skips its terminal hook at an actionable-wait
1765
2224
  // boundary, so its `finally` routes the disconnect here — and
@@ -1768,9 +2227,21 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1768
2227
  // still threw on the next request. An explicit cancel is different: the
1769
2228
  // user gave up on the approval, so the cancel band stays authoritative.
1770
2229
  const state = runState.get(ctx)
1771
- if (cancelled || (!detachableRun(ctx) && state?.interrupted !== true)) {
1772
- await abortRun(runs, ctx.runId, state?.usage)
1773
- return
2230
+ let terminal = false
2231
+ try {
2232
+ // The durable cancel read is best-effort. It must not bypass the
2233
+ // terminal persistence path or prevent the completion promise from
2234
+ // settling when the run store is unavailable.
2235
+ const cancelled =
2236
+ info.cancelRequested === true ||
2237
+ (runs !== undefined && (await wasCancelRequested(runs, ctx.runId)))
2238
+ terminal =
2239
+ cancelled || (!detachableRun(ctx) && state?.interrupted !== true)
2240
+ if (terminal) {
2241
+ await abortRun(runs, ctx.runId, state?.usage)
2242
+ }
2243
+ } finally {
2244
+ if (terminal) state?.completion?.reject(info.reason)
1774
2245
  }
1775
2246
  // A plain disconnect on a detachable or interrupted run: write NOTHING.
1776
2247
  // Either the agent is still running and a later attach can take it over