@tanstack/ai-persistence 0.1.4 → 0.2.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.
@@ -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
@@ -35,6 +35,7 @@ import type {
35
35
  RunAgentResumeItem,
36
36
  StreamChunk,
37
37
  ToolApprovalResolution,
38
+ BilledUsage,
38
39
  TokenUsage,
39
40
  } from '@tanstack/ai'
40
41
  import type {
@@ -254,6 +255,8 @@ interface RunStateEntry {
254
255
  pending: Array<InterruptRecord>
255
256
  resumeByInterruptId: Map<string, RunAgentResumeItem>
256
257
  }
258
+ /** Usage accumulated across every model call in this chat invocation. */
259
+ usage?: TokenUsage
257
260
  /** Accumulated terminal-turn text, for throttled streaming snapshots (B). */
258
261
  streamingText?: string
259
262
  /** Epoch ms of the last streaming snapshot, to throttle writes (B). */
@@ -1287,12 +1290,98 @@ async function createOrResumeRun(
1287
1290
  runs: RunStore | undefined,
1288
1291
  runId: string,
1289
1292
  threadId: string,
1290
- ): Promise<void> {
1291
- await runs?.createOrResume({
1293
+ ): Promise<TokenUsage | undefined> {
1294
+ const run = await runs?.createOrResume({
1292
1295
  runId,
1293
1296
  threadId,
1294
1297
  startedAt: Date.now(),
1295
1298
  })
1299
+ return run?.usage
1300
+ }
1301
+
1302
+ function sumOptionalNumber(
1303
+ current: number | undefined,
1304
+ next: number | undefined,
1305
+ ): number | undefined {
1306
+ if (current === undefined) return next
1307
+ if (next === undefined) return current
1308
+ return current + next
1309
+ }
1310
+
1311
+ function sumNumberFields<T extends object>(
1312
+ current: T | undefined,
1313
+ next: T | undefined,
1314
+ ): T | undefined {
1315
+ if (!current) return next
1316
+ if (!next) return current
1317
+
1318
+ const result = { ...current }
1319
+ for (const key of Object.keys(next) as Array<keyof T>) {
1320
+ const currentValue = current[key]
1321
+ const nextValue = next[key]
1322
+ if (typeof nextValue === 'number') {
1323
+ result[key] = ((typeof currentValue === 'number' ? currentValue : 0) +
1324
+ nextValue) as T[keyof T]
1325
+ }
1326
+ }
1327
+ return result
1328
+ }
1329
+
1330
+ function accumulateTokenUsage(
1331
+ current: TokenUsage | undefined,
1332
+ next: TokenUsage,
1333
+ ): TokenUsage {
1334
+ if (!current) return { ...next }
1335
+
1336
+ const promptTokensDetails = sumNumberFields(
1337
+ current.promptTokensDetails,
1338
+ next.promptTokensDetails,
1339
+ )
1340
+ const completionTokensDetails = sumNumberFields(
1341
+ current.completionTokensDetails,
1342
+ next.completionTokensDetails,
1343
+ )
1344
+ const costDetails = sumNumberFields(current.costDetails, next.costDetails)
1345
+ // Provider-specific details are opaque, so retain the latest reported bag.
1346
+ const providerUsageDetails =
1347
+ next.providerUsageDetails ?? current.providerUsageDetails
1348
+ const durationSeconds = sumOptionalNumber(
1349
+ current.durationSeconds,
1350
+ next.durationSeconds,
1351
+ )
1352
+ const unitsBilled = sumOptionalNumber(current.unitsBilled, next.unitsBilled)
1353
+ const billed = accumulateBilled(current.billed, next.billed)
1354
+ const cost = sumOptionalNumber(current.cost, next.cost)
1355
+
1356
+ return {
1357
+ ...current,
1358
+ ...next,
1359
+ promptTokens: current.promptTokens + next.promptTokens,
1360
+ completionTokens: current.completionTokens + next.completionTokens,
1361
+ totalTokens: current.totalTokens + next.totalTokens,
1362
+ ...(promptTokensDetails ? { promptTokensDetails } : {}),
1363
+ ...(completionTokensDetails ? { completionTokensDetails } : {}),
1364
+ ...(durationSeconds !== undefined ? { durationSeconds } : {}),
1365
+ ...(unitsBilled !== undefined ? { unitsBilled } : {}),
1366
+ ...(billed !== undefined ? { billed } : {}),
1367
+ ...(cost !== undefined ? { cost } : {}),
1368
+ ...(costDetails ? { costDetails } : {}),
1369
+ ...(providerUsageDetails ? { providerUsageDetails } : {}),
1370
+ }
1371
+ }
1372
+
1373
+ /**
1374
+ * Sum billed quantities when both reports use the same unit. Different units
1375
+ * cannot be added, so the later report wins.
1376
+ */
1377
+ function accumulateBilled(
1378
+ current: BilledUsage | undefined,
1379
+ next: BilledUsage | undefined,
1380
+ ): BilledUsage | undefined {
1381
+ if (!current) return next
1382
+ if (!next) return current
1383
+ if (current.unit !== next.unit) return next
1384
+ return { quantity: current.quantity + next.quantity, unit: current.unit }
1296
1385
  }
1297
1386
 
1298
1387
  async function completeRun(
@@ -1311,6 +1400,7 @@ async function failRun(
1311
1400
  runs: RunStore | undefined,
1312
1401
  runId: string,
1313
1402
  error: unknown,
1403
+ usage?: TokenUsage,
1314
1404
  ): Promise<void> {
1315
1405
  // `RunRecord.error` is a structured `RunError`. Only `message` is filled in
1316
1406
  // here: the middleware sees an opaque thrown value, and inventing a `code`
@@ -1320,6 +1410,7 @@ async function failRun(
1320
1410
  status: 'failed',
1321
1411
  finishedAt: Date.now(),
1322
1412
  error: { message: error instanceof Error ? error.message : String(error) },
1413
+ ...(usage ? { usage } : {}),
1323
1414
  })
1324
1415
  }
1325
1416
 
@@ -1334,9 +1425,11 @@ async function failRun(
1334
1425
  export async function interruptRun(
1335
1426
  runs: RunStore | undefined,
1336
1427
  runId: string,
1428
+ usage?: TokenUsage,
1337
1429
  ): Promise<void> {
1338
1430
  await runs?.update(runId, {
1339
1431
  status: 'interrupted',
1432
+ ...(usage ? { usage } : {}),
1340
1433
  })
1341
1434
  }
1342
1435
 
@@ -1348,10 +1441,12 @@ export async function interruptRun(
1348
1441
  export async function abortRun(
1349
1442
  runs: RunStore | undefined,
1350
1443
  runId: string,
1444
+ usage?: TokenUsage,
1351
1445
  ): Promise<void> {
1352
1446
  await runs?.update(runId, {
1353
1447
  status: 'aborted',
1354
1448
  finishedAt: Date.now(),
1449
+ ...(usage ? { usage } : {}),
1355
1450
  })
1356
1451
  }
1357
1452
 
@@ -1508,15 +1603,15 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1508
1603
  }
1509
1604
  }
1510
1605
 
1511
- await createOrResumeRun(runs, ctx.runId, ctx.threadId)
1606
+ const storedUsage = await createOrResumeRun(runs, ctx.runId, ctx.threadId)
1512
1607
 
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
- }
1608
+ const state = runState.get(ctx)
1609
+ // A continuation has a fresh middleware context but resumes the same run.
1610
+ if (state && storedUsage) state.usage = storedUsage
1611
+ if (!state?.merged) {
1612
+ if (state) state.merged = true
1613
+ const stored = await messageStore.loadThread(ctx.threadId)
1614
+ patch.messages = config.messages.length > 0 ? config.messages : stored
1520
1615
  }
1521
1616
 
1522
1617
  return Object.keys(patch).length > 0 ? patch : undefined
@@ -1542,7 +1637,14 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1542
1637
  if (ctx.phase === 'modelStream') {
1543
1638
  const s = runState.get(ctx)
1544
1639
  if (s && chunk.type === 'TEXT_MESSAGE_START') {
1545
- s.streamingMessageId = chunk.messageId
1640
+ // An empty/malformed messageId means "no identity" (matching the
1641
+ // engine's convention), leaving room for the TOOL_CALL_START
1642
+ // parentMessageId fallback below — but the per-turn accumulator
1643
+ // still resets so snapshots never mix text across turns.
1644
+ s.streamingMessageId =
1645
+ typeof chunk.messageId === 'string' && chunk.messageId !== ''
1646
+ ? chunk.messageId
1647
+ : undefined
1546
1648
  s.streamingMessageCreatedAt = new Date()
1547
1649
  s.streamingText = ''
1548
1650
  } else if (
@@ -1622,11 +1724,24 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1622
1724
  })
1623
1725
  }
1624
1726
  }
1625
- await interruptRun(runs, ctx.runId)
1727
+ // Adapter terminals arrive before `onUsage`; synthesized tool boundaries
1728
+ // arrive after it with the same usage already in state.
1729
+ const usage =
1730
+ ctx.phase === 'modelStream' && chunk.usage
1731
+ ? accumulateTokenUsage(state.usage, chunk.usage)
1732
+ : (state.usage ?? chunk.usage)
1733
+ state.usage = usage
1734
+ await interruptRun(runs, ctx.runId, usage)
1626
1735
  await messageStore.saveThread(ctx.threadId, [...ctx.messages])
1627
1736
  state.interrupted = true
1628
1737
  },
1629
1738
 
1739
+ onUsage(ctx: ChatMiddlewareContext, usage: TokenUsage) {
1740
+ const state = runState.get(ctx)
1741
+ if (!state || state.interrupted) return
1742
+ state.usage = accumulateTokenUsage(state.usage, usage)
1743
+ },
1744
+
1630
1745
  async onFinish(ctx: ChatMiddlewareContext, info: FinishInfo) {
1631
1746
  const state = runState.get(ctx)
1632
1747
  if (state?.interrupted) return
@@ -1643,12 +1758,12 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1643
1758
  state?.streamingMessageCreatedAt,
1644
1759
  ),
1645
1760
  )
1646
- await completeRun(runs, ctx.runId, info.usage)
1761
+ await completeRun(runs, ctx.runId, state?.usage ?? info.usage)
1647
1762
  await commitPendingResumes(state, persistence.stores.interrupts)
1648
1763
  },
1649
1764
 
1650
1765
  async onError(ctx: ChatMiddlewareContext, info: ErrorInfo) {
1651
- await failRun(runs, ctx.runId, info.error)
1766
+ await failRun(runs, ctx.runId, info.error, runState.get(ctx)?.usage)
1652
1767
  },
1653
1768
 
1654
1769
  async onAbort(ctx: ChatMiddlewareContext, info: AbortInfo) {
@@ -1671,7 +1786,7 @@ export function withPersistence<TStores extends ChatTranscriptStores>(
1671
1786
  // user gave up on the approval, so the cancel band stays authoritative.
1672
1787
  const state = runState.get(ctx)
1673
1788
  if (cancelled || (!detachableRun(ctx) && state?.interrupted !== true)) {
1674
- await abortRun(runs, ctx.runId)
1789
+ await abortRun(runs, ctx.runId, state?.usage)
1675
1790
  return
1676
1791
  }
1677
1792
  // 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.