UNPKG

@tanstack/ai-persistence

Version:

Composable state persistence for TanStack AI messages, runs, interrupts, metadata, and locks.

360 lines (359 loc) 12.1 kB
import { defineAIPersistence } from "./types.js"; import { resolveBlobRange } from "./blob-range.js"; //#region src/memory.ts var compareUtf8Bytes = (left, right) => { const leftBytes = new TextEncoder().encode(left); const rightBytes = new TextEncoder().encode(right); const length = Math.min(leftBytes.length, rightBytes.length); for (let index = 0; index < length; index++) { const leftByte = leftBytes[index]; const rightByte = rightBytes[index]; if (leftByte !== rightByte) return (leftByte ?? 0) - (rightByte ?? 0); } return leftBytes.length - rightBytes.length; }; var MemoryMessageStore = class { threads = /* @__PURE__ */ new Map(); loadThread(threadId) { return Promise.resolve(this.threads.get(threadId)?.slice() ?? []); } saveThread(threadId, messages) { this.threads.set(threadId, messages.slice()); return Promise.resolve(); } }; var MemoryRunStore = class { runs = /* @__PURE__ */ new Map(); createOrResume(input) { const existing = this.runs.get(input.runId); if (existing) return Promise.resolve(existing); const record = { runId: input.runId, threadId: input.threadId, status: input.status ?? "running", startedAt: input.startedAt }; this.runs.set(record.runId, record); return Promise.resolve(record); } update(runId, patch) { const existing = this.runs.get(runId); if (existing) this.runs.set(runId, { ...existing, ...patch }); return Promise.resolve(); } get(runId) { return Promise.resolve(this.runs.get(runId) ?? null); } findActiveRun(threadId) { const active = [...this.runs.values()].filter((run) => run.threadId === threadId && run.status === "running").sort((a, b) => b.startedAt - a.startedAt); return Promise.resolve(active[0] ?? null); } listByThread(threadId) { const matching = [...this.runs.values()].filter((run) => run.threadId === threadId).sort((a, b) => a.startedAt - b.startedAt); return Promise.resolve(matching); } listReclaimable(opts) { const cutoff = opts.now - opts.ttlMs; const matching = [...this.runs.values()].filter((run) => run.status === "running" && run.detachedSince !== void 0 && run.detachedSince <= cutoff); return Promise.resolve(matching); } }; var MemoryGenerationRunStore = class { generationRuns = /* @__PURE__ */ new Map(); createOrResume(input) { const existing = this.generationRuns.get(input.runId); if (existing) return Promise.resolve(existing); const record = { runId: input.runId, threadId: input.threadId, activity: input.activity, provider: input.provider, model: input.model, status: input.status ?? "running", startedAt: input.startedAt }; this.generationRuns.set(record.runId, record); return Promise.resolve(record); } update(runId, patch) { const existing = this.generationRuns.get(runId); if (existing) this.generationRuns.set(runId, { ...existing, ...patch }); return Promise.resolve(); } get(runId) { return Promise.resolve(this.generationRuns.get(runId) ?? null); } findLatestForThread(threadId) { const linked = [...this.generationRuns.values()].filter((run) => run.threadId === threadId).sort((a, b) => b.startedAt - a.startedAt); return Promise.resolve(linked[0] ?? null); } }; function byRequestedAt(a, b) { return a.requestedAt - b.requestedAt; } var MemoryInterruptStore = class { interrupts = /* @__PURE__ */ new Map(); create(record) { if (!this.interrupts.has(record.interruptId)) this.interrupts.set(record.interruptId, { ...record, status: "pending" }); return Promise.resolve(); } resolve(interruptId, response) { const existing = this.interrupts.get(interruptId); if (existing) this.interrupts.set(interruptId, { ...existing, status: "resolved", resolvedAt: Date.now(), response }); return Promise.resolve(); } cancel(interruptId) { const existing = this.interrupts.get(interruptId); if (existing) this.interrupts.set(interruptId, { ...existing, status: "cancelled", resolvedAt: Date.now() }); return Promise.resolve(); } async commitBatch(entries) { const ids = /* @__PURE__ */ new Set(); for (const entry of entries) { if (ids.has(entry.interruptId)) throw new Error(`Interrupt batch contains duplicate id: ${entry.interruptId}.`); ids.add(entry.interruptId); const existing = this.interrupts.get(entry.interruptId); if (!existing) throw new Error(`Interrupt batch references missing id: ${entry.interruptId}.`); if (existing.status !== "pending") throw new Error(`Interrupt batch references non-pending id: ${entry.interruptId}.`); } const resolvedAt = Date.now(); for (const entry of entries) { const existing = this.interrupts.get(entry.interruptId); if (!existing) continue; if (entry.status === "resolved") this.interrupts.set(entry.interruptId, { ...existing, status: "resolved", resolvedAt, response: entry.response }); else this.interrupts.set(entry.interruptId, { ...existing, status: "cancelled", resolvedAt }); } } get(interruptId) { return Promise.resolve(this.interrupts.get(interruptId) ?? null); } list(threadId) { return Promise.resolve([...this.interrupts.values()].filter((interrupt) => interrupt.threadId === threadId).sort(byRequestedAt)); } listPending(threadId) { return Promise.resolve([...this.interrupts.values()].filter((interrupt) => interrupt.threadId === threadId && interrupt.status === "pending").sort(byRequestedAt)); } listByRun(runId) { return Promise.resolve([...this.interrupts.values()].filter((interrupt) => interrupt.runId === runId).sort(byRequestedAt)); } listPendingByRun(runId) { return Promise.resolve([...this.interrupts.values()].filter((interrupt) => interrupt.runId === runId && interrupt.status === "pending").sort(byRequestedAt)); } }; var MemoryMetadataStore = class { values = /* @__PURE__ */ new Map(); get(namespace, key) { const bucket = this.values.get(namespace); if (!bucket || !bucket.has(key)) return Promise.resolve(null); return Promise.resolve(bucket.get(key)); } set(namespace, key, value) { let bucket = this.values.get(namespace); if (!bucket) { bucket = /* @__PURE__ */ new Map(); this.values.set(namespace, bucket); } bucket.set(key, value); return Promise.resolve(); } delete(namespace, key) { const bucket = this.values.get(namespace); if (!bucket) return Promise.resolve(); bucket.delete(key); if (bucket.size === 0) this.values.delete(namespace); return Promise.resolve(); } }; var MemoryArtifactStore = class { artifacts = /* @__PURE__ */ new Map(); save(record) { this.artifacts.set(record.artifactId, { ...record }); return Promise.resolve(); } get(artifactId) { return Promise.resolve(this.artifacts.get(artifactId) ?? null); } list(runId) { return Promise.resolve([...this.artifacts.values()].filter((a) => a.runId === runId).sort((a, b) => a.createdAt - b.createdAt || compareUtf8Bytes(a.artifactId, b.artifactId))); } listForThread(threadId) { return Promise.resolve([...this.artifacts.values()].filter((a) => a.threadId === threadId).sort((a, b) => a.createdAt - b.createdAt || compareUtf8Bytes(a.artifactId, b.artifactId))); } delete(artifactId) { this.artifacts.delete(artifactId); return Promise.resolve(); } deleteForRun(runId) { for (const artifact of this.artifacts.values()) if (artifact.runId === runId) this.artifacts.delete(artifact.artifactId); return Promise.resolve(); } }; var textEncoder = new TextEncoder(); var textDecoder = new TextDecoder(); function copyBytes(bytes) { return new Uint8Array(bytes); } function bytesToArrayBuffer(bytes) { const buffer = new ArrayBuffer(bytes.byteLength); new Uint8Array(buffer).set(bytes); return buffer; } async function bytesFromStream(stream) { const reader = stream.getReader(); const chunks = []; let total = 0; try { while (true) { const { done, value } = await reader.read(); if (done) break; chunks.push(copyBytes(value)); total += value.byteLength; } } finally { reader.releaseLock(); } const bytes = new Uint8Array(total); let offset = 0; for (const chunk of chunks) { bytes.set(chunk, offset); offset += chunk.byteLength; } return bytes; } async function bytesFromBlobBody(body) { if (typeof body === "string") return textEncoder.encode(body); if (body instanceof ArrayBuffer) return new Uint8Array(body.slice(0)); if (ArrayBuffer.isView(body)) return copyBytes(new Uint8Array(body.buffer, body.byteOffset, body.byteLength)); if (typeof Blob !== "undefined" && body instanceof Blob) return new Uint8Array(await body.arrayBuffer()); if (typeof ReadableStream !== "undefined" && body instanceof ReadableStream) return bytesFromStream(body); throw new TypeError("Unsupported blob body."); } function blobRecordSnapshot(record) { return { ...record, ...record.customMetadata ? { customMetadata: { ...record.customMetadata } } : {} }; } function blobObject(record, bytes, range) { const served = range ? resolveBlobRange(bytes.byteLength, range) : { offset: 0, length: bytes.byteLength }; const copied = copyBytes(bytes.subarray(served.offset, served.offset + served.length)); return { ...blobRecordSnapshot(record), ...range ? { range: served } : {}, body: new ReadableStream({ start(controller) { controller.enqueue(copyBytes(copied)); controller.close(); } }), arrayBuffer: () => Promise.resolve(bytesToArrayBuffer(copied)), text: () => Promise.resolve(textDecoder.decode(copied)) }; } var MemoryBlobStore = class { blobs = /* @__PURE__ */ new Map(); nextEtag = 1; async put(key, body, options) { const bytes = await bytesFromBlobBody(body); const existing = this.blobs.get(key); const now = Date.now(); const record = { key, size: bytes.byteLength, etag: String(this.nextEtag++), contentType: options?.contentType ?? (typeof Blob !== "undefined" && body instanceof Blob ? body.type || void 0 : void 0), customMetadata: options?.customMetadata ? { ...options.customMetadata } : void 0, createdAt: existing?.record.createdAt ?? now, updatedAt: now }; this.blobs.set(key, { record, bytes: copyBytes(bytes) }); return blobRecordSnapshot(record); } get(key, options) { const entry = this.blobs.get(key); return Promise.resolve(entry ? blobObject(entry.record, entry.bytes, options?.range) : null); } head(key) { const entry = this.blobs.get(key); return Promise.resolve(entry ? blobRecordSnapshot(entry.record) : null); } delete(key) { this.blobs.delete(key); return Promise.resolve(); } list(options) { const limit = options?.limit; if (limit === 0) return Promise.resolve({ objects: [], truncated: false }); const keys = [...this.blobs.keys()].filter((key) => key.startsWith(options?.prefix ?? "")).filter((key) => options?.cursor === void 0 || key > options.cursor).sort(); const pageKeys = limit === void 0 ? keys : keys.slice(0, limit); const objects = pageKeys.map((key) => { const blob = this.blobs.get(key); if (blob === void 0) throw new Error(`Missing blob for listed key: ${key}`); return blobRecordSnapshot(blob.record); }); const truncated = limit !== void 0 && keys.length > limit; return Promise.resolve({ objects, ...truncated ? { cursor: pageKeys.at(-1), truncated } : {} }); } }; /** * In-process reference backend for the full state + generation store set. * * Returns messages + runs + generationRuns + interrupts + metadata + artifacts * + blobs. Locks are not included — use `InMemoryLockStore` + `withLocks` from * `@tanstack/ai` when a test or single-process app needs coordination. */ function memoryPersistence() { const stores = { messages: new MemoryMessageStore(), runs: new MemoryRunStore(), generationRuns: new MemoryGenerationRunStore(), interrupts: new MemoryInterruptStore(), metadata: new MemoryMetadataStore(), artifacts: new MemoryArtifactStore(), blobs: new MemoryBlobStore() }; return defineAIPersistence({ stores }); } //#endregion export { memoryPersistence }; //# sourceMappingURL=memory.js.map