@tanstack/ai-persistence 0.1.3 → 0.1.5

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.
@@ -172,6 +172,11 @@ predicate: `(status: RunStatus) => status is TerminalRunStatus`, so calling it
172
172
  inside a guard narrows `status` to `TerminalRunStatus` for the rest of that
173
173
  branch, with no cast needed.
174
174
 
175
+ `RunRecord.usage` is optional. `withPersistence` sums reported numeric fields
176
+ across provider calls for that `runId`, while opaque `providerUsageDetails`
177
+ retains the latest reported bag. Known usage is persisted on interruption and
178
+ every terminal status.
179
+
175
180
  `RunRecord.error` is a structured `RunError`, not a bare string:
176
181
 
177
182
  ```ts
@@ -253,10 +258,10 @@ store through `update`/`get` — but `cancelRequested` must round-trip
253
258
  faithfully (previous section) for the durable path to work at all.
254
259
 
255
260
  - **`createOrResume`** (required): if `runId` exists, return it **unchanged**,
256
- ignoring the passed `threadId` / `startedAt` / `status`. Resuming a run does
257
- not reset `startedAt` or overwrite its current status. Idempotent retries and
258
- double-submit depend on this. `status` defaults to `'running'` on first
259
- creation.
261
+ including its stored `usage`, and ignore the passed `threadId` / `startedAt` /
262
+ `status`. Resuming a run does not reset `startedAt` or overwrite its current
263
+ status. Idempotent retries and double-submit depend on this. `status` defaults
264
+ to `'running'` on first creation.
260
265
  - **`update`** (required): missing `runId` is a **no-op** (do not throw, do not
261
266
  insert).
262
267
  - **`get`** (required): current record, or `null` when unknown.
package/src/middleware.ts CHANGED
@@ -254,6 +254,8 @@ interface RunStateEntry {
254
254
  pending: Array<InterruptRecord>
255
255
  resumeByInterruptId: Map<string, RunAgentResumeItem>
256
256
  }
257
+ /** Usage accumulated across every model call in this chat invocation. */
258
+ usage?: TokenUsage
257
259
  /** Accumulated terminal-turn text, for throttled streaming snapshots (B). */
258
260
  streamingText?: string
259
261
  /** Epoch ms of the last streaming snapshot, to throttle writes (B). */
@@ -265,6 +267,7 @@ interface RunStateEntry {
265
267
  * bubble in place.
266
268
  */
267
269
  streamingMessageId?: string
270
+ streamingMessageCreatedAt?: Date
268
271
  }
269
272
 
270
273
  const runState = new WeakMap<object, RunStateEntry>()
@@ -450,6 +453,7 @@ function finishedTranscript(
450
453
  messages: ReadonlyArray<ModelMessage>,
451
454
  info: FinishInfo,
452
455
  messageId: string | undefined,
456
+ createdAt: Date | undefined,
453
457
  ): Array<ModelMessage> {
454
458
  const transcript = [...messages]
455
459
  const last = transcript[transcript.length - 1]
@@ -464,6 +468,7 @@ function finishedTranscript(
464
468
  role: 'assistant',
465
469
  content: info.content,
466
470
  ...(messageId ? { id: messageId } : {}),
471
+ ...(createdAt ? { createdAt } : {}),
467
472
  })
468
473
  }
469
474
  return transcript
@@ -1284,12 +1289,82 @@ async function createOrResumeRun(
1284
1289
  runs: RunStore | undefined,
1285
1290
  runId: string,
1286
1291
  threadId: string,
1287
- ): Promise<void> {
1288
- await runs?.createOrResume({
1292
+ ): Promise<TokenUsage | undefined> {
1293
+ const run = await runs?.createOrResume({
1289
1294
  runId,
1290
1295
  threadId,
1291
1296
  startedAt: Date.now(),
1292
1297
  })
1298
+ return run?.usage
1299
+ }
1300
+
1301
+ function sumOptionalNumber(
1302
+ current: number | undefined,
1303
+ next: number | undefined,
1304
+ ): number | undefined {
1305
+ if (current === undefined) return next
1306
+ if (next === undefined) return current
1307
+ return current + next
1308
+ }
1309
+
1310
+ function sumNumberFields<T extends object>(
1311
+ current: T | undefined,
1312
+ next: T | undefined,
1313
+ ): T | undefined {
1314
+ if (!current) return next
1315
+ if (!next) return current
1316
+
1317
+ const result = { ...current }
1318
+ for (const key of Object.keys(next) as Array<keyof T>) {
1319
+ const currentValue = current[key]
1320
+ const nextValue = next[key]
1321
+ if (typeof nextValue === 'number') {
1322
+ result[key] = ((typeof currentValue === 'number' ? currentValue : 0) +
1323
+ nextValue) as T[keyof T]
1324
+ }
1325
+ }
1326
+ return result
1327
+ }
1328
+
1329
+ function accumulateTokenUsage(
1330
+ current: TokenUsage | undefined,
1331
+ next: TokenUsage,
1332
+ ): TokenUsage {
1333
+ if (!current) return { ...next }
1334
+
1335
+ const promptTokensDetails = sumNumberFields(
1336
+ current.promptTokensDetails,
1337
+ next.promptTokensDetails,
1338
+ )
1339
+ const completionTokensDetails = sumNumberFields(
1340
+ current.completionTokensDetails,
1341
+ next.completionTokensDetails,
1342
+ )
1343
+ const costDetails = sumNumberFields(current.costDetails, next.costDetails)
1344
+ // Provider-specific details are opaque, so retain the latest reported bag.
1345
+ const providerUsageDetails =
1346
+ next.providerUsageDetails ?? current.providerUsageDetails
1347
+ const durationSeconds = sumOptionalNumber(
1348
+ current.durationSeconds,
1349
+ next.durationSeconds,
1350
+ )
1351
+ const unitsBilled = sumOptionalNumber(current.unitsBilled, next.unitsBilled)
1352
+ const cost = sumOptionalNumber(current.cost, next.cost)
1353
+
1354
+ return {
1355
+ ...current,
1356
+ ...next,
1357
+ promptTokens: current.promptTokens + next.promptTokens,
1358
+ completionTokens: current.completionTokens + next.completionTokens,
1359
+ totalTokens: current.totalTokens + next.totalTokens,
1360
+ ...(promptTokensDetails ? { promptTokensDetails } : {}),
1361
+ ...(completionTokensDetails ? { completionTokensDetails } : {}),
1362
+ ...(durationSeconds !== undefined ? { durationSeconds } : {}),
1363
+ ...(unitsBilled !== undefined ? { unitsBilled } : {}),
1364
+ ...(cost !== undefined ? { cost } : {}),
1365
+ ...(costDetails ? { costDetails } : {}),
1366
+ ...(providerUsageDetails ? { providerUsageDetails } : {}),
1367
+ }
1293
1368
  }
1294
1369
 
1295
1370
  async function completeRun(
@@ -1308,6 +1383,7 @@ async function failRun(
1308
1383
  runs: RunStore | undefined,
1309
1384
  runId: string,
1310
1385
  error: unknown,
1386
+ usage?: TokenUsage,
1311
1387
  ): Promise<void> {
1312
1388
  // `RunRecord.error` is a structured `RunError`. Only `message` is filled in
1313
1389
  // here: the middleware sees an opaque thrown value, and inventing a `code`
@@ -1317,6 +1393,7 @@ async function failRun(
1317
1393
  status: 'failed',
1318
1394
  finishedAt: Date.now(),
1319
1395
  error: { message: error instanceof Error ? error.message : String(error) },
1396
+ ...(usage ? { usage } : {}),
1320
1397
  })
1321
1398
  }
1322
1399
 
@@ -1331,9 +1408,11 @@ async function failRun(
1331
1408
  export async function interruptRun(
1332
1409
  runs: RunStore | undefined,
1333
1410
  runId: string,
1411
+ usage?: TokenUsage,
1334
1412
  ): Promise<void> {
1335
1413
  await runs?.update(runId, {
1336
1414
  status: 'interrupted',
1415
+ ...(usage ? { usage } : {}),
1337
1416
  })
1338
1417
  }
1339
1418
 
@@ -1345,10 +1424,12 @@ export async function interruptRun(
1345
1424
  export async function abortRun(
1346
1425
  runs: RunStore | undefined,
1347
1426
  runId: string,
1427
+ usage?: TokenUsage,
1348
1428
  ): Promise<void> {
1349
1429
  await runs?.update(runId, {
1350
1430
  status: 'aborted',
1351
1431
  finishedAt: Date.now(),
1432
+ ...(usage ? { usage } : {}),
1352
1433
  })
1353
1434
  }
1354
1435
 
@@ -1505,15 +1586,15 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1505
1586
  }
1506
1587
  }
1507
1588
 
1508
- await createOrResumeRun(runs, ctx.runId, ctx.threadId)
1589
+ const storedUsage = await createOrResumeRun(runs, ctx.runId, ctx.threadId)
1509
1590
 
1510
- {
1511
- const state = runState.get(ctx)
1512
- if (!state?.merged) {
1513
- if (state) state.merged = true
1514
- const stored = await messageStore.loadThread(ctx.threadId)
1515
- patch.messages = config.messages.length > 0 ? config.messages : stored
1516
- }
1591
+ const state = runState.get(ctx)
1592
+ // A continuation has a fresh middleware context but resumes the same run.
1593
+ if (state && storedUsage) state.usage = storedUsage
1594
+ if (!state?.merged) {
1595
+ if (state) state.merged = true
1596
+ const stored = await messageStore.loadThread(ctx.threadId)
1597
+ patch.messages = config.messages.length > 0 ? config.messages : stored
1517
1598
  }
1518
1599
 
1519
1600
  return Object.keys(patch).length > 0 ? patch : undefined
@@ -1536,11 +1617,28 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1536
1617
  // regardless of snapshotStreaming — it's persisted onto the assistant
1537
1618
  // message so its identity survives hydrate and a reload resumes the same
1538
1619
  // bubble in place.
1539
- if (chunk.type === 'TEXT_MESSAGE_START') {
1620
+ if (ctx.phase === 'modelStream') {
1540
1621
  const s = runState.get(ctx)
1541
- if (s) {
1542
- s.streamingMessageId = chunk.messageId
1622
+ if (s && chunk.type === 'TEXT_MESSAGE_START') {
1623
+ // An empty/malformed messageId means "no identity" (matching the
1624
+ // engine's convention), leaving room for the TOOL_CALL_START
1625
+ // parentMessageId fallback below — but the per-turn accumulator
1626
+ // still resets so snapshots never mix text across turns.
1627
+ s.streamingMessageId =
1628
+ typeof chunk.messageId === 'string' && chunk.messageId !== ''
1629
+ ? chunk.messageId
1630
+ : undefined
1631
+ s.streamingMessageCreatedAt = new Date()
1543
1632
  s.streamingText = ''
1633
+ } else if (
1634
+ s &&
1635
+ chunk.type === 'TOOL_CALL_START' &&
1636
+ typeof chunk.parentMessageId === 'string' &&
1637
+ chunk.parentMessageId !== '' &&
1638
+ s.streamingMessageId === undefined
1639
+ ) {
1640
+ s.streamingMessageId = chunk.parentMessageId
1641
+ s.streamingMessageCreatedAt ??= new Date()
1544
1642
  }
1545
1643
  }
1546
1644
 
@@ -1571,6 +1669,9 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1571
1669
  ...(snapshotState.streamingMessageId
1572
1670
  ? { id: snapshotState.streamingMessageId }
1573
1671
  : {}),
1672
+ ...(snapshotState.streamingMessageCreatedAt
1673
+ ? { createdAt: snapshotState.streamingMessageCreatedAt }
1674
+ : {}),
1574
1675
  },
1575
1676
  ])
1576
1677
  } catch {
@@ -1606,11 +1707,24 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1606
1707
  })
1607
1708
  }
1608
1709
  }
1609
- await interruptRun(runs, ctx.runId)
1710
+ // Adapter terminals arrive before `onUsage`; synthesized tool boundaries
1711
+ // arrive after it with the same usage already in state.
1712
+ const usage =
1713
+ ctx.phase === 'modelStream' && chunk.usage
1714
+ ? accumulateTokenUsage(state.usage, chunk.usage)
1715
+ : (state.usage ?? chunk.usage)
1716
+ state.usage = usage
1717
+ await interruptRun(runs, ctx.runId, usage)
1610
1718
  await messageStore.saveThread(ctx.threadId, [...ctx.messages])
1611
1719
  state.interrupted = true
1612
1720
  },
1613
1721
 
1722
+ onUsage(ctx: ChatMiddlewareContext, usage: TokenUsage) {
1723
+ const state = runState.get(ctx)
1724
+ if (!state || state.interrupted) return
1725
+ state.usage = accumulateTokenUsage(state.usage, usage)
1726
+ },
1727
+
1614
1728
  async onFinish(ctx: ChatMiddlewareContext, info: FinishInfo) {
1615
1729
  const state = runState.get(ctx)
1616
1730
  if (state?.interrupted) return
@@ -1620,14 +1734,19 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1620
1734
  // "finished" run whose transcript is missing the terminal turn.
1621
1735
  await messageStore.saveThread(
1622
1736
  ctx.threadId,
1623
- finishedTranscript(ctx.messages, info, state?.streamingMessageId),
1737
+ finishedTranscript(
1738
+ ctx.messages,
1739
+ info,
1740
+ state?.streamingMessageId,
1741
+ state?.streamingMessageCreatedAt,
1742
+ ),
1624
1743
  )
1625
- await completeRun(runs, ctx.runId, info.usage)
1744
+ await completeRun(runs, ctx.runId, state?.usage ?? info.usage)
1626
1745
  await commitPendingResumes(state, persistence.stores.interrupts)
1627
1746
  },
1628
1747
 
1629
1748
  async onError(ctx: ChatMiddlewareContext, info: ErrorInfo) {
1630
- await failRun(runs, ctx.runId, info.error)
1749
+ await failRun(runs, ctx.runId, info.error, runState.get(ctx)?.usage)
1631
1750
  },
1632
1751
 
1633
1752
  async onAbort(ctx: ChatMiddlewareContext, info: AbortInfo) {
@@ -1650,7 +1769,7 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1650
1769
  // user gave up on the approval, so the cancel band stays authoritative.
1651
1770
  const state = runState.get(ctx)
1652
1771
  if (cancelled || (!detachableRun(ctx) && state?.interrupted !== true)) {
1653
- await abortRun(runs, ctx.runId)
1772
+ await abortRun(runs, ctx.runId, state?.usage)
1654
1773
  return
1655
1774
  }
1656
1775
  // A plain disconnect on a detachable or interrupted run: write NOTHING.
@@ -299,6 +299,13 @@ export function runPersistenceConformance(
299
299
  usage: { promptTokens: 3, completionTokens: 4, totalTokens: 7 },
300
300
  })
301
301
 
302
+ const resumedAfterUpdate = await store.createOrResume({
303
+ runId: 'run-1',
304
+ threadId: 'thread-different',
305
+ startedAt: 9999,
306
+ })
307
+ expect(resumedAfterUpdate).toEqual(done)
308
+
302
309
  // `error` is a structured RunError: the prose `message` plus the
303
310
  // optional machine-branchable `code`. Both must survive the round-trip,
304
311
  // so a backend that flattens the record to a bare string fails here.