@tanstack/ai-persistence 0.1.4 → 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). */
@@ -1287,12 +1289,82 @@ async function createOrResumeRun(
1287
1289
  runs: RunStore | undefined,
1288
1290
  runId: string,
1289
1291
  threadId: string,
1290
- ): Promise<void> {
1291
- await runs?.createOrResume({
1292
+ ): Promise<TokenUsage | undefined> {
1293
+ const run = await runs?.createOrResume({
1292
1294
  runId,
1293
1295
  threadId,
1294
1296
  startedAt: Date.now(),
1295
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
+ }
1296
1368
  }
1297
1369
 
1298
1370
  async function completeRun(
@@ -1311,6 +1383,7 @@ async function failRun(
1311
1383
  runs: RunStore | undefined,
1312
1384
  runId: string,
1313
1385
  error: unknown,
1386
+ usage?: TokenUsage,
1314
1387
  ): Promise<void> {
1315
1388
  // `RunRecord.error` is a structured `RunError`. Only `message` is filled in
1316
1389
  // here: the middleware sees an opaque thrown value, and inventing a `code`
@@ -1320,6 +1393,7 @@ async function failRun(
1320
1393
  status: 'failed',
1321
1394
  finishedAt: Date.now(),
1322
1395
  error: { message: error instanceof Error ? error.message : String(error) },
1396
+ ...(usage ? { usage } : {}),
1323
1397
  })
1324
1398
  }
1325
1399
 
@@ -1334,9 +1408,11 @@ async function failRun(
1334
1408
  export async function interruptRun(
1335
1409
  runs: RunStore | undefined,
1336
1410
  runId: string,
1411
+ usage?: TokenUsage,
1337
1412
  ): Promise<void> {
1338
1413
  await runs?.update(runId, {
1339
1414
  status: 'interrupted',
1415
+ ...(usage ? { usage } : {}),
1340
1416
  })
1341
1417
  }
1342
1418
 
@@ -1348,10 +1424,12 @@ export async function interruptRun(
1348
1424
  export async function abortRun(
1349
1425
  runs: RunStore | undefined,
1350
1426
  runId: string,
1427
+ usage?: TokenUsage,
1351
1428
  ): Promise<void> {
1352
1429
  await runs?.update(runId, {
1353
1430
  status: 'aborted',
1354
1431
  finishedAt: Date.now(),
1432
+ ...(usage ? { usage } : {}),
1355
1433
  })
1356
1434
  }
1357
1435
 
@@ -1508,15 +1586,15 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1508
1586
  }
1509
1587
  }
1510
1588
 
1511
- await createOrResumeRun(runs, ctx.runId, ctx.threadId)
1589
+ const storedUsage = await createOrResumeRun(runs, ctx.runId, ctx.threadId)
1512
1590
 
1513
- {
1514
- const state = runState.get(ctx)
1515
- if (!state?.merged) {
1516
- if (state) state.merged = true
1517
- const stored = await messageStore.loadThread(ctx.threadId)
1518
- patch.messages = config.messages.length > 0 ? config.messages : stored
1519
- }
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
1520
1598
  }
1521
1599
 
1522
1600
  return Object.keys(patch).length > 0 ? patch : undefined
@@ -1542,7 +1620,14 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1542
1620
  if (ctx.phase === 'modelStream') {
1543
1621
  const s = runState.get(ctx)
1544
1622
  if (s && chunk.type === 'TEXT_MESSAGE_START') {
1545
- s.streamingMessageId = chunk.messageId
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
1546
1631
  s.streamingMessageCreatedAt = new Date()
1547
1632
  s.streamingText = ''
1548
1633
  } else if (
@@ -1622,11 +1707,24 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1622
1707
  })
1623
1708
  }
1624
1709
  }
1625
- 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)
1626
1718
  await messageStore.saveThread(ctx.threadId, [...ctx.messages])
1627
1719
  state.interrupted = true
1628
1720
  },
1629
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
+
1630
1728
  async onFinish(ctx: ChatMiddlewareContext, info: FinishInfo) {
1631
1729
  const state = runState.get(ctx)
1632
1730
  if (state?.interrupted) return
@@ -1643,12 +1741,12 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1643
1741
  state?.streamingMessageCreatedAt,
1644
1742
  ),
1645
1743
  )
1646
- await completeRun(runs, ctx.runId, info.usage)
1744
+ await completeRun(runs, ctx.runId, state?.usage ?? info.usage)
1647
1745
  await commitPendingResumes(state, persistence.stores.interrupts)
1648
1746
  },
1649
1747
 
1650
1748
  async onError(ctx: ChatMiddlewareContext, info: ErrorInfo) {
1651
- await failRun(runs, ctx.runId, info.error)
1749
+ await failRun(runs, ctx.runId, info.error, runState.get(ctx)?.usage)
1652
1750
  },
1653
1751
 
1654
1752
  async onAbort(ctx: ChatMiddlewareContext, info: AbortInfo) {
@@ -1671,7 +1769,7 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1671
1769
  // user gave up on the approval, so the cancel band stays authoritative.
1672
1770
  const state = runState.get(ctx)
1673
1771
  if (cancelled || (!detachableRun(ctx) && state?.interrupted !== true)) {
1674
- await abortRun(runs, ctx.runId)
1772
+ await abortRun(runs, ctx.runId, state?.usage)
1675
1773
  return
1676
1774
  }
1677
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.