@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.
- package/dist/esm/memory.js +3 -2
- package/dist/esm/memory.js.map +1 -1
- package/dist/esm/middleware.d.ts +3 -3
- package/dist/esm/middleware.js +90 -23
- package/dist/esm/middleware.js.map +1 -1
- package/dist/esm/testkit/conformance.js +52 -29
- package/dist/esm/testkit/conformance.js.map +1 -1
- package/package.json +4 -4
- package/skills/ai-persistence/stores/SKILL.md +9 -4
- package/src/middleware.ts +130 -15
- package/src/testkit/conformance.ts +7 -0
|
@@ -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
|
-
|
|
257
|
-
not reset `startedAt` or overwrite its current
|
|
258
|
-
double-submit depend on this. `status` defaults
|
|
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<
|
|
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
|
-
|
|
1515
|
-
|
|
1516
|
-
|
|
1517
|
-
|
|
1518
|
-
|
|
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
|
-
|
|
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
|
-
|
|
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.
|