@tanstack/ai-persistence
Version:
Composable state persistence for TanStack AI messages, runs, interrupts, metadata, and locks.
360 lines (359 loc) • 12.1 kB
JavaScript
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