@tanstack/ai-persistence 0.0.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/blob-range.d.ts +51 -0
- package/dist/esm/blob-range.js +84 -0
- package/dist/esm/blob-range.js.map +1 -0
- package/dist/esm/capabilities.d.ts +5 -0
- package/dist/esm/capabilities.js +16 -0
- package/dist/esm/capabilities.js.map +1 -0
- package/dist/esm/index.d.ts +13 -0
- package/dist/esm/index.js +9 -0
- package/dist/esm/memory.d.ts +19 -0
- package/dist/esm/memory.js +319 -0
- package/dist/esm/memory.js.map +1 -0
- package/dist/esm/middleware.d.ts +252 -0
- package/dist/esm/middleware.js +872 -0
- package/dist/esm/middleware.js.map +1 -0
- package/dist/esm/reconstruct-generation.d.ts +129 -0
- package/dist/esm/reconstruct-generation.js +148 -0
- package/dist/esm/reconstruct-generation.js.map +1 -0
- package/dist/esm/reconstruct.d.ts +79 -0
- package/dist/esm/reconstruct.js +75 -0
- package/dist/esm/reconstruct.js.map +1 -0
- package/dist/esm/retrieve.d.ts +40 -0
- package/dist/esm/retrieve.js +54 -0
- package/dist/esm/retrieve.js.map +1 -0
- package/dist/esm/testkit/conformance.d.ts +33 -0
- package/dist/esm/testkit/conformance.js +997 -0
- package/dist/esm/testkit/conformance.js.map +1 -0
- package/dist/esm/types.d.ts +554 -0
- package/dist/esm/types.js +103 -0
- package/dist/esm/types.js.map +1 -0
- package/package.json +71 -0
- package/skills/ai-persistence/SKILL.md +218 -0
- package/skills/ai-persistence/build-cloudflare-adapter/SKILL.md +313 -0
- package/skills/ai-persistence/build-cloudflare-artifact-store/SKILL.md +693 -0
- package/skills/ai-persistence/build-custom-adapter/SKILL.md +328 -0
- package/skills/ai-persistence/build-drizzle-adapter/SKILL.md +562 -0
- package/skills/ai-persistence/build-prisma-adapter/SKILL.md +518 -0
- package/skills/ai-persistence/server/SKILL.md +210 -0
- package/skills/ai-persistence/stores/SKILL.md +485 -0
- package/src/blob-range.ts +101 -0
- package/src/capabilities.ts +18 -0
- package/src/index.ts +114 -0
- package/src/memory.ts +491 -0
- package/src/middleware.ts +1795 -0
- package/src/reconstruct-generation.ts +244 -0
- package/src/reconstruct.ts +149 -0
- package/src/retrieve.ts +77 -0
- package/src/testkit/conformance.ts +1288 -0
- package/src/types.ts +878 -0
|
@@ -0,0 +1,872 @@
|
|
|
1
|
+
import { validateChatPersistenceStores, validateGenerationPersistenceStores } from "./types.js";
|
|
2
|
+
import { InterruptsCapability, PersistenceCapability, provideInterrupts, providePersistence } from "./capabilities.js";
|
|
3
|
+
import { artifactBlobKey } from "./retrieve.js";
|
|
4
|
+
import { defineChatMiddleware, getDetachableRun, wasCancelRequested } from "@tanstack/ai";
|
|
5
|
+
import { providePendingTurn } from "@tanstack/ai/adapter-internals";
|
|
6
|
+
import { base64ToUint8Array } from "@tanstack/ai-utils";
|
|
7
|
+
//#region src/middleware.ts
|
|
8
|
+
/**
|
|
9
|
+
* The slot this generation's runs are filed under: `ctx.threadId` (the
|
|
10
|
+
* `threadId` the caller passed the activity), or the option when it overrides.
|
|
11
|
+
*
|
|
12
|
+
* Throws when neither supplies one. A run filed under no scope can never be
|
|
13
|
+
* hydrated by one, so `persistence: true` would restore nothing, forever. That
|
|
14
|
+
* is worth failing loudly for, since the alternative is a silent hole a reader
|
|
15
|
+
* cannot diagnose from behavior.
|
|
16
|
+
*/
|
|
17
|
+
function generationScope(ctx, opts) {
|
|
18
|
+
const threadId = opts.threadId ?? ctx.threadId;
|
|
19
|
+
if (threadId === void 0 || threadId.length === 0) throw new Error("Generation persistence requires a `threadId`, the stable scope successive runs are filed under. Pass it to the activity, e.g. `generateImage({ threadId, middleware: [withGenerationPersistence(p)] })`, or override it with `withGenerationPersistence(p, { threadId })`.");
|
|
20
|
+
return threadId;
|
|
21
|
+
}
|
|
22
|
+
var DEFAULT_ARTIFACT_FETCH_TIMEOUT_MS = 3e4;
|
|
23
|
+
var DEFAULT_MAX_ARTIFACT_BYTES = 1024 * 1024 * 1024;
|
|
24
|
+
var runState = /* @__PURE__ */ new WeakMap();
|
|
25
|
+
var validResumeStatuses = /* @__PURE__ */ new Set(["resolved", "cancelled"]);
|
|
26
|
+
function validatePendingResumes(pending, resume) {
|
|
27
|
+
const pendingInterruptIds = new Set(pending.map((interrupt) => interrupt.interruptId));
|
|
28
|
+
const resumeByInterruptId = new Map((resume ?? []).map((entry) => [entry.interruptId, entry]));
|
|
29
|
+
if (pending.length === 0) {
|
|
30
|
+
const staleEntry = resume?.[0];
|
|
31
|
+
if (staleEntry) throw new Error(`Resume entry references non-pending interrupt ${staleEntry.interruptId}.`);
|
|
32
|
+
return resumeByInterruptId;
|
|
33
|
+
}
|
|
34
|
+
if (!resume || resume.length === 0) throw new Error(`Thread has pending interrupts; resume is required before accepting new input.`);
|
|
35
|
+
for (const interrupt of pending) {
|
|
36
|
+
const entry = resumeByInterruptId.get(interrupt.interruptId);
|
|
37
|
+
if (!entry) throw new Error(`Missing resume entry for pending interrupt ${interrupt.interruptId}.`);
|
|
38
|
+
if (!validResumeStatuses.has(entry.status)) throw new Error(`Invalid resume status for pending interrupt ${interrupt.interruptId}: ${entry.status}.`);
|
|
39
|
+
}
|
|
40
|
+
for (const entry of resume) if (!pendingInterruptIds.has(entry.interruptId)) throw new Error(`Resume entry references non-pending interrupt ${entry.interruptId}.`);
|
|
41
|
+
return resumeByInterruptId;
|
|
42
|
+
}
|
|
43
|
+
async function applyPendingResumes(pending, resumeByInterruptId, interrupts) {
|
|
44
|
+
for (const interrupt of pending) {
|
|
45
|
+
const entry = resumeByInterruptId.get(interrupt.interruptId);
|
|
46
|
+
if (!entry) continue;
|
|
47
|
+
if (entry.status === "resolved") await interrupts.resolve(interrupt.interruptId, entry.payload);
|
|
48
|
+
else await interrupts.cancel(interrupt.interruptId);
|
|
49
|
+
}
|
|
50
|
+
}
|
|
51
|
+
/**
|
|
52
|
+
* Commit the resumes stashed in `onConfig`, marking each resumed interrupt
|
|
53
|
+
* resolved/cancelled. Called only from success boundaries (`onFinish`, and the
|
|
54
|
+
* `onChunk` interrupt boundary) so a provider failure or abort between accepting
|
|
55
|
+
* the resume and reaching a boundary leaves the interrupts pending — the
|
|
56
|
+
* approval is not consumed and a retry with the same resume succeeds. Idempotent
|
|
57
|
+
* and a no-op when nothing is stashed.
|
|
58
|
+
*/
|
|
59
|
+
async function commitPendingResumes(state, interrupts) {
|
|
60
|
+
if (!state?.pendingResumes || !interrupts) return;
|
|
61
|
+
const { pending, resumeByInterruptId } = state.pendingResumes;
|
|
62
|
+
await applyPendingResumes(pending, resumeByInterruptId, interrupts);
|
|
63
|
+
state.pendingResumes = void 0;
|
|
64
|
+
}
|
|
65
|
+
function objectValue(value) {
|
|
66
|
+
return value && typeof value === "object" ? value : null;
|
|
67
|
+
}
|
|
68
|
+
function stringField(value, key) {
|
|
69
|
+
return typeof value[key] === "string" ? value[key] : void 0;
|
|
70
|
+
}
|
|
71
|
+
function interruptKind(interrupt) {
|
|
72
|
+
const metadata = objectValue(interrupt.payload.metadata);
|
|
73
|
+
return metadata ? stringField(metadata, "kind") : void 0;
|
|
74
|
+
}
|
|
75
|
+
function resolvedApprovalDecision(entry) {
|
|
76
|
+
if (entry.status === "cancelled") return false;
|
|
77
|
+
const payload = objectValue(entry.payload);
|
|
78
|
+
return typeof payload?.approved === "boolean" ? payload.approved : false;
|
|
79
|
+
}
|
|
80
|
+
/**
|
|
81
|
+
* Translate the persisted pending interrupts + the resume batch into the
|
|
82
|
+
* `ChatResumeToolState` the chat engine consumes. This is the server-authoritative
|
|
83
|
+
* counterpart to the engine's ephemeral (client-history) reconstruction: because
|
|
84
|
+
* the persistence flow sends empty client messages, the engine has no history to
|
|
85
|
+
* rebuild from, so persistence supplies the resume state directly (and clears
|
|
86
|
+
* `config.resume` so the ephemeral path is skipped — see `onConfig`).
|
|
87
|
+
*/
|
|
88
|
+
function resumeToolStateFromPending(pending, resumeByInterruptId) {
|
|
89
|
+
const approvals = /* @__PURE__ */ new Map();
|
|
90
|
+
const clientToolResults = /* @__PURE__ */ new Map();
|
|
91
|
+
for (const interrupt of pending) {
|
|
92
|
+
const entry = resumeByInterruptId.get(interrupt.interruptId);
|
|
93
|
+
if (!entry) continue;
|
|
94
|
+
const kind = interruptKind(interrupt);
|
|
95
|
+
const reason = stringField(interrupt.payload, "reason");
|
|
96
|
+
const toolCallId = stringField(interrupt.payload, "toolCallId");
|
|
97
|
+
if (kind === "approval" || reason === "approval_required") {
|
|
98
|
+
approvals.set(interrupt.interruptId, resolvedApprovalDecision(entry));
|
|
99
|
+
continue;
|
|
100
|
+
}
|
|
101
|
+
if (entry.status === "resolved" && toolCallId && (kind === "client_tool" || reason === "client_tool_input")) clientToolResults.set(toolCallId, entry.payload);
|
|
102
|
+
}
|
|
103
|
+
if (approvals.size === 0 && clientToolResults.size === 0) return void 0;
|
|
104
|
+
return {
|
|
105
|
+
approvals,
|
|
106
|
+
clientToolResults
|
|
107
|
+
};
|
|
108
|
+
}
|
|
109
|
+
/**
|
|
110
|
+
* Build the transcript to persist when a run finishes successfully.
|
|
111
|
+
*
|
|
112
|
+
* The chat engine appends an assistant message to the middleware message list
|
|
113
|
+
* only when that turn carries tool calls (to feed the agent loop); a run's
|
|
114
|
+
* terminal *text* reply is never appended. So `ctx.messages` at `onFinish` is
|
|
115
|
+
* missing the assistant's final answer. Reattach it from the finish info —
|
|
116
|
+
* `info.content` is the last turn's accumulated text (reset each cycle) — so
|
|
117
|
+
* the stored thread is the complete conversation a server-authoritative client
|
|
118
|
+
* hydrates on load. A guard avoids duplicating a terminal assistant turn should
|
|
119
|
+
* the engine ever start appending it itself.
|
|
120
|
+
*/
|
|
121
|
+
function finishedTranscript(messages, info, messageId) {
|
|
122
|
+
const transcript = [...messages];
|
|
123
|
+
const last = transcript[transcript.length - 1];
|
|
124
|
+
const alreadyPresent = last?.role === "assistant" && last.toolCalls === void 0 && last.content === info.content;
|
|
125
|
+
if (info.content && !alreadyPresent) transcript.push({
|
|
126
|
+
role: "assistant",
|
|
127
|
+
content: info.content,
|
|
128
|
+
...messageId ? { id: messageId } : {}
|
|
129
|
+
});
|
|
130
|
+
return transcript;
|
|
131
|
+
}
|
|
132
|
+
function interruptPayload(interrupt) {
|
|
133
|
+
return interrupt && typeof interrupt === "object" ? { ...interrupt } : { value: interrupt };
|
|
134
|
+
}
|
|
135
|
+
function isArtifactRef(value) {
|
|
136
|
+
const record = objectValue(value);
|
|
137
|
+
return !!record && typeof record.artifactId === "string";
|
|
138
|
+
}
|
|
139
|
+
function mediaActivity(activity) {
|
|
140
|
+
return activity === "image" || activity === "audio" || activity === "tts" || activity === "video" || activity === "transcription" ? activity : void 0;
|
|
141
|
+
}
|
|
142
|
+
function parseDataUrl(value) {
|
|
143
|
+
const match = /^data:([^;,]+)?(;base64)?,(.*)$/s.exec(value);
|
|
144
|
+
if (!match) return void 0;
|
|
145
|
+
const mimeType = match[1] || "application/octet-stream";
|
|
146
|
+
const raw = match[3] ?? "";
|
|
147
|
+
let payload;
|
|
148
|
+
try {
|
|
149
|
+
payload = decodeURIComponent(raw);
|
|
150
|
+
} catch {
|
|
151
|
+
payload = raw;
|
|
152
|
+
}
|
|
153
|
+
return {
|
|
154
|
+
mimeType,
|
|
155
|
+
bytes: match[2] ? base64ToUint8Array(payload) : new TextEncoder().encode(payload)
|
|
156
|
+
};
|
|
157
|
+
}
|
|
158
|
+
function extensionForMime(mimeType) {
|
|
159
|
+
if (mimeType === void 0) return "bin";
|
|
160
|
+
switch (mimeType) {
|
|
161
|
+
case "image/png": return "png";
|
|
162
|
+
case "image/jpeg": return "jpg";
|
|
163
|
+
case "audio/wav": return "wav";
|
|
164
|
+
case "audio/mpeg": return "mp3";
|
|
165
|
+
case "audio/mp3": return "mp3";
|
|
166
|
+
case "video/mp4": return "mp4";
|
|
167
|
+
case "application/json": return "json";
|
|
168
|
+
default: return "bin";
|
|
169
|
+
}
|
|
170
|
+
}
|
|
171
|
+
function defaultArtifactName(descriptor, activity, index) {
|
|
172
|
+
const ext = extensionForMime(descriptor.mimeType);
|
|
173
|
+
return `${activity}-${descriptor.role}-${descriptor.mediaType ?? "artifact"}-${index}.${ext}`;
|
|
174
|
+
}
|
|
175
|
+
function sourcePartDescriptors(part, role, path) {
|
|
176
|
+
const record = objectValue(part);
|
|
177
|
+
const type = stringField(record ?? {}, "type");
|
|
178
|
+
const source = objectValue(record?.source);
|
|
179
|
+
if (!record || !source || type !== "image" && type !== "audio" && type !== "video") return [];
|
|
180
|
+
const sourceType = stringField(source, "type");
|
|
181
|
+
const mimeType = stringField(source, "mimeType") ?? `${type}/mpeg`;
|
|
182
|
+
if (sourceType === "data") {
|
|
183
|
+
const value = stringField(source, "value");
|
|
184
|
+
if (!value) return [];
|
|
185
|
+
return [{
|
|
186
|
+
role,
|
|
187
|
+
path,
|
|
188
|
+
mediaType: type,
|
|
189
|
+
mimeType,
|
|
190
|
+
bytes: base64ToUint8Array(value)
|
|
191
|
+
}];
|
|
192
|
+
}
|
|
193
|
+
if (sourceType === "url") {
|
|
194
|
+
const value = stringField(source, "value");
|
|
195
|
+
if (!value) return [];
|
|
196
|
+
return [{
|
|
197
|
+
role,
|
|
198
|
+
path,
|
|
199
|
+
mediaType: type,
|
|
200
|
+
mimeType,
|
|
201
|
+
url: value
|
|
202
|
+
}];
|
|
203
|
+
}
|
|
204
|
+
return [];
|
|
205
|
+
}
|
|
206
|
+
function promptInputDescriptors(inputs) {
|
|
207
|
+
const prompt = objectValue(inputs)?.prompt;
|
|
208
|
+
if (!Array.isArray(prompt)) return [];
|
|
209
|
+
const counts = {
|
|
210
|
+
image: 0,
|
|
211
|
+
audio: 0,
|
|
212
|
+
video: 0
|
|
213
|
+
};
|
|
214
|
+
const descriptors = [];
|
|
215
|
+
for (const part of prompt) {
|
|
216
|
+
const type = stringField(objectValue(part) ?? {}, "type");
|
|
217
|
+
if (type !== "image" && type !== "audio" && type !== "video") continue;
|
|
218
|
+
const index = counts[type] ?? 0;
|
|
219
|
+
counts[type] = index + 1;
|
|
220
|
+
descriptors.push(...sourcePartDescriptors(part, "input", `prompt.${type}s.${index}`));
|
|
221
|
+
}
|
|
222
|
+
return descriptors;
|
|
223
|
+
}
|
|
224
|
+
function generatedMediaDescriptor(args) {
|
|
225
|
+
const media = objectValue(args.media);
|
|
226
|
+
if (!media) return void 0;
|
|
227
|
+
const b64Json = stringField(media, "b64Json");
|
|
228
|
+
if (b64Json) return {
|
|
229
|
+
role: args.role,
|
|
230
|
+
path: args.path,
|
|
231
|
+
mediaType: args.mediaType,
|
|
232
|
+
mimeType: stringField(media, "contentType") ?? args.mimeType,
|
|
233
|
+
bytes: base64ToUint8Array(b64Json),
|
|
234
|
+
jobId: args.jobId,
|
|
235
|
+
expiresAt: args.expiresAt
|
|
236
|
+
};
|
|
237
|
+
const url = stringField(media, "url");
|
|
238
|
+
if (url) return {
|
|
239
|
+
role: args.role,
|
|
240
|
+
path: args.path,
|
|
241
|
+
mediaType: args.mediaType,
|
|
242
|
+
mimeType: stringField(media, "contentType") ?? args.mimeType,
|
|
243
|
+
url,
|
|
244
|
+
jobId: args.jobId,
|
|
245
|
+
expiresAt: args.expiresAt
|
|
246
|
+
};
|
|
247
|
+
}
|
|
248
|
+
function builtInArtifactDescriptors(activity, inputs, result) {
|
|
249
|
+
const descriptors = promptInputDescriptors(inputs);
|
|
250
|
+
const output = objectValue(result);
|
|
251
|
+
if (!output) return descriptors;
|
|
252
|
+
if (activity === "image" && Array.isArray(output.images)) output.images.forEach((image, index) => {
|
|
253
|
+
const descriptor = generatedMediaDescriptor({
|
|
254
|
+
role: "output",
|
|
255
|
+
path: `images.${index}`,
|
|
256
|
+
mediaType: "image",
|
|
257
|
+
mimeType: "image/png",
|
|
258
|
+
media: image
|
|
259
|
+
});
|
|
260
|
+
if (descriptor) descriptors.push(descriptor);
|
|
261
|
+
});
|
|
262
|
+
if (activity === "audio") {
|
|
263
|
+
const descriptor = generatedMediaDescriptor({
|
|
264
|
+
role: "output",
|
|
265
|
+
path: "audio",
|
|
266
|
+
mediaType: "audio",
|
|
267
|
+
mimeType: "audio/mpeg",
|
|
268
|
+
media: output.audio
|
|
269
|
+
});
|
|
270
|
+
if (descriptor) descriptors.push(descriptor);
|
|
271
|
+
}
|
|
272
|
+
if (activity === "tts") {
|
|
273
|
+
const audio = stringField(output, "audio");
|
|
274
|
+
if (audio) {
|
|
275
|
+
const format = stringField(output, "format");
|
|
276
|
+
descriptors.push({
|
|
277
|
+
role: "output",
|
|
278
|
+
path: "audio",
|
|
279
|
+
mediaType: "audio",
|
|
280
|
+
mimeType: stringField(output, "contentType") ?? (format ? `audio/${format}` : "audio/mpeg"),
|
|
281
|
+
bytes: base64ToUint8Array(audio)
|
|
282
|
+
});
|
|
283
|
+
}
|
|
284
|
+
}
|
|
285
|
+
if (activity === "video" && typeof output.url === "string") descriptors.push({
|
|
286
|
+
role: "output",
|
|
287
|
+
path: "video",
|
|
288
|
+
mediaType: "video",
|
|
289
|
+
mimeType: "video/mp4",
|
|
290
|
+
url: output.url,
|
|
291
|
+
jobId: stringField(output, "jobId"),
|
|
292
|
+
expiresAt: output.expiresAt instanceof Date ? output.expiresAt : void 0
|
|
293
|
+
});
|
|
294
|
+
if (activity === "transcription") {
|
|
295
|
+
const audio = objectValue(inputs)?.audio;
|
|
296
|
+
if (typeof audio === "string") {
|
|
297
|
+
const data = parseDataUrl(audio);
|
|
298
|
+
descriptors.push({
|
|
299
|
+
role: "input",
|
|
300
|
+
path: "audio",
|
|
301
|
+
mediaType: "audio",
|
|
302
|
+
mimeType: data?.mimeType ?? "audio/mpeg",
|
|
303
|
+
bytes: data?.bytes ?? base64ToUint8Array(audio)
|
|
304
|
+
});
|
|
305
|
+
} else if (audio instanceof ArrayBuffer) descriptors.push({
|
|
306
|
+
role: "input",
|
|
307
|
+
path: "audio",
|
|
308
|
+
mediaType: "audio",
|
|
309
|
+
mimeType: "audio/mpeg",
|
|
310
|
+
bytes: audio.slice(0)
|
|
311
|
+
});
|
|
312
|
+
else if (typeof Blob !== "undefined" && audio instanceof Blob) descriptors.push({
|
|
313
|
+
role: "input",
|
|
314
|
+
path: "audio",
|
|
315
|
+
mediaType: "audio",
|
|
316
|
+
mimeType: audio.type || "audio/mpeg",
|
|
317
|
+
bytes: audio
|
|
318
|
+
});
|
|
319
|
+
if (Array.isArray(output.segments) || Array.isArray(output.words)) descriptors.push({
|
|
320
|
+
role: "output",
|
|
321
|
+
path: "transcription",
|
|
322
|
+
mediaType: "json",
|
|
323
|
+
mimeType: "application/json",
|
|
324
|
+
json: output
|
|
325
|
+
});
|
|
326
|
+
}
|
|
327
|
+
return descriptors;
|
|
328
|
+
}
|
|
329
|
+
/**
|
|
330
|
+
* Reject hosts that only make sense as an SSRF target: loopback, link-local
|
|
331
|
+
* (including the cloud metadata address), private, and unique-local ranges.
|
|
332
|
+
*
|
|
333
|
+
* Applied to caller-supplied input URLs only. Provider result URLs skip it on
|
|
334
|
+
* purpose — a self-hosted or local provider legitimately returns a `localhost`
|
|
335
|
+
* URL, and those live inside the same trust boundary as the adapter itself.
|
|
336
|
+
*
|
|
337
|
+
* This checks IP *literals*. A hostname that resolves to a private address
|
|
338
|
+
* passes, which is why `allowInputUrl` is required rather than optional.
|
|
339
|
+
*/
|
|
340
|
+
function isBlockedInputHost(hostname) {
|
|
341
|
+
const host = hostname.toLowerCase().replace(/^\[|\]$/g, "");
|
|
342
|
+
if (host === "localhost" || host.endsWith(".localhost")) return true;
|
|
343
|
+
const ipv4 = /^(\d{1,3})\.(\d{1,3})\.(\d{1,3})\.(\d{1,3})$/.exec(host);
|
|
344
|
+
if (ipv4) {
|
|
345
|
+
const [a, b] = [Number(ipv4[1]), Number(ipv4[2])];
|
|
346
|
+
if (a === 127 || a === 0 || a === 10) return true;
|
|
347
|
+
if (a === 169 && b === 254) return true;
|
|
348
|
+
if (a === 172 && b >= 16 && b <= 31) return true;
|
|
349
|
+
if (a === 192 && b === 168) return true;
|
|
350
|
+
return false;
|
|
351
|
+
}
|
|
352
|
+
if (host === "::" || host === "::1") return true;
|
|
353
|
+
if (host.startsWith("fe80:")) return true;
|
|
354
|
+
if (/^f[cd][0-9a-f]{2}:/.test(host)) return true;
|
|
355
|
+
const mappedDotted = /^::ffff:(\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3})$/.exec(host);
|
|
356
|
+
if (mappedDotted?.[1]) return isBlockedInputHost(mappedDotted[1]);
|
|
357
|
+
const mappedHex = /^::ffff:([0-9a-f]{1,4}):([0-9a-f]{1,4})$/.exec(host);
|
|
358
|
+
if (mappedHex?.[1] && mappedHex[2]) {
|
|
359
|
+
const high = Number.parseInt(mappedHex[1], 16);
|
|
360
|
+
const low = Number.parseInt(mappedHex[2], 16);
|
|
361
|
+
return isBlockedInputHost(`${high >> 8}.${high & 255}.${low >> 8}.${low & 255}`);
|
|
362
|
+
}
|
|
363
|
+
return false;
|
|
364
|
+
}
|
|
365
|
+
/**
|
|
366
|
+
* Fail the stream once more than `maxBytes` have passed through, so an
|
|
367
|
+
* unexpectedly huge artifact can't fill the blob store.
|
|
368
|
+
*
|
|
369
|
+
* Only used when the response does NOT already bound itself — a chunked reply,
|
|
370
|
+
* or a content-encoded one whose declared length describes the compressed
|
|
371
|
+
* bytes. When `content-length` describes the body the store will drain, HTTP
|
|
372
|
+
* framing is the bound and wrapping would only cost the caller the declared
|
|
373
|
+
* length: a `TransformStream`'s readable side carries none, which is what
|
|
374
|
+
* pushes a length-strict runtime (workerd + R2) onto a multipart upload.
|
|
375
|
+
*/
|
|
376
|
+
function capBodySize(body, maxBytes, url) {
|
|
377
|
+
let seen = 0;
|
|
378
|
+
return body.pipeThrough(new TransformStream({ transform(chunk, controller) {
|
|
379
|
+
seen += chunk.byteLength;
|
|
380
|
+
if (seen > maxBytes) {
|
|
381
|
+
controller.error(/* @__PURE__ */ new Error(`Artifact at ${url} exceeds maxArtifactBytes (${maxBytes}).`));
|
|
382
|
+
return;
|
|
383
|
+
}
|
|
384
|
+
controller.enqueue(chunk);
|
|
385
|
+
} }));
|
|
386
|
+
}
|
|
387
|
+
/**
|
|
388
|
+
* Resolve a descriptor to the bytes to store. Returns `undefined` when the
|
|
389
|
+
* descriptor is deliberately not persisted — today that means a caller-supplied
|
|
390
|
+
* input URL without an `allowInputUrl` opt-in.
|
|
391
|
+
*/
|
|
392
|
+
async function descriptorBody(descriptor, opts) {
|
|
393
|
+
if (descriptor.json !== void 0) {
|
|
394
|
+
const body = JSON.stringify(descriptor.json);
|
|
395
|
+
return {
|
|
396
|
+
body,
|
|
397
|
+
size: new TextEncoder().encode(body).byteLength,
|
|
398
|
+
mimeType: descriptor.mimeType ?? "application/json"
|
|
399
|
+
};
|
|
400
|
+
}
|
|
401
|
+
if (descriptor.bytes !== void 0) {
|
|
402
|
+
const body = descriptor.bytes;
|
|
403
|
+
let size;
|
|
404
|
+
if (typeof body === "string") size = new TextEncoder().encode(body).byteLength;
|
|
405
|
+
else if (body instanceof ArrayBuffer) size = body.byteLength;
|
|
406
|
+
else if (ArrayBuffer.isView(body)) size = body.byteLength;
|
|
407
|
+
else if (typeof Blob !== "undefined" && body instanceof Blob) size = body.size;
|
|
408
|
+
else size = 0;
|
|
409
|
+
return {
|
|
410
|
+
body,
|
|
411
|
+
size,
|
|
412
|
+
mimeType: descriptor.mimeType ?? "application/octet-stream"
|
|
413
|
+
};
|
|
414
|
+
}
|
|
415
|
+
if (descriptor.url) {
|
|
416
|
+
const data = parseDataUrl(descriptor.url);
|
|
417
|
+
if (data) return {
|
|
418
|
+
body: data.bytes,
|
|
419
|
+
size: data.bytes.byteLength,
|
|
420
|
+
mimeType: descriptor.mimeType ?? data.mimeType
|
|
421
|
+
};
|
|
422
|
+
const isCallerSupplied = descriptor.role === "input";
|
|
423
|
+
const allowInputUrl = opts?.allowInputUrl;
|
|
424
|
+
if (isCallerSupplied && !allowInputUrl) return void 0;
|
|
425
|
+
let target;
|
|
426
|
+
try {
|
|
427
|
+
target = new URL(descriptor.url);
|
|
428
|
+
} catch {
|
|
429
|
+
throw new Error(`Failed to persist artifact: ${descriptor.url} is not a valid URL.`);
|
|
430
|
+
}
|
|
431
|
+
if (target.protocol !== "https:" && target.protocol !== "http:") throw new Error(`Refusing to fetch artifact over ${target.protocol} (${descriptor.path}).`);
|
|
432
|
+
if (allowInputUrl && isCallerSupplied) {
|
|
433
|
+
if (isBlockedInputHost(target.hostname)) throw new Error(`Refusing to fetch input artifact from internal host ${target.hostname}.`);
|
|
434
|
+
if (!await allowInputUrl({
|
|
435
|
+
url: target,
|
|
436
|
+
descriptor
|
|
437
|
+
})) throw new Error(`Refusing to fetch input artifact from ${target.hostname}: rejected by allowInputUrl.`);
|
|
438
|
+
}
|
|
439
|
+
const maxBytes = opts?.maxArtifactBytes ?? DEFAULT_MAX_ARTIFACT_BYTES;
|
|
440
|
+
const response = await (opts?.artifactFetch ?? globalThis.fetch)(target, {
|
|
441
|
+
redirect: isCallerSupplied ? "manual" : "follow",
|
|
442
|
+
signal: AbortSignal.timeout(opts?.artifactFetchTimeoutMs ?? DEFAULT_ARTIFACT_FETCH_TIMEOUT_MS)
|
|
443
|
+
});
|
|
444
|
+
if (isCallerSupplied && response.status >= 300 && response.status < 400) throw new Error(`Refusing to follow a redirect for input artifact ${descriptor.path}.`);
|
|
445
|
+
if (!response.ok) throw new Error(`Failed to persist artifact from ${descriptor.url}: HTTP ${response.status}`);
|
|
446
|
+
const contentLength = response.headers.get("content-length");
|
|
447
|
+
const declaredLength = contentLength === null ? void 0 : Number(contentLength);
|
|
448
|
+
if (maxBytes !== false && declaredLength !== void 0 && Number.isFinite(declaredLength) && declaredLength > maxBytes) throw new Error(`Artifact at ${descriptor.url} exceeds maxArtifactBytes (${maxBytes}).`);
|
|
449
|
+
const mimeType = descriptor.mimeType ?? response.headers.get("content-type") ?? "application/octet-stream";
|
|
450
|
+
const encoding = response.headers.get("content-encoding");
|
|
451
|
+
const decodedLengthIsKnown = declaredLength !== void 0 && Number.isFinite(declaredLength) && (encoding === null || encoding === "identity");
|
|
452
|
+
const expectedLength = decodedLengthIsKnown ? declaredLength : void 0;
|
|
453
|
+
if (response.body) return {
|
|
454
|
+
body: maxBytes === false || decodedLengthIsKnown ? response.body : capBodySize(response.body, maxBytes, descriptor.url),
|
|
455
|
+
size: 0,
|
|
456
|
+
expectedLength,
|
|
457
|
+
mimeType,
|
|
458
|
+
sourceUrl: descriptor.url
|
|
459
|
+
};
|
|
460
|
+
const body = await response.arrayBuffer();
|
|
461
|
+
if (maxBytes !== false && body.byteLength > maxBytes) throw new Error(`Artifact at ${descriptor.url} exceeds maxArtifactBytes (${maxBytes}).`);
|
|
462
|
+
return {
|
|
463
|
+
body,
|
|
464
|
+
size: body.byteLength,
|
|
465
|
+
mimeType,
|
|
466
|
+
sourceUrl: descriptor.url
|
|
467
|
+
};
|
|
468
|
+
}
|
|
469
|
+
throw new Error(`Artifact descriptor ${descriptor.path} has no bytes, url, or json.`);
|
|
470
|
+
}
|
|
471
|
+
async function persistGenerationArtifacts(persistence, opts, ctx, result) {
|
|
472
|
+
const activity = mediaActivity(ctx.activity);
|
|
473
|
+
if (!activity) return [];
|
|
474
|
+
const threadId = generationScope(ctx, opts);
|
|
475
|
+
const runId = ctx.runId ?? ctx.requestId;
|
|
476
|
+
const extractionInput = {
|
|
477
|
+
activity,
|
|
478
|
+
provider: ctx.provider,
|
|
479
|
+
model: ctx.model,
|
|
480
|
+
threadId,
|
|
481
|
+
runId,
|
|
482
|
+
inputs: ctx.artifactInputs,
|
|
483
|
+
result
|
|
484
|
+
};
|
|
485
|
+
const extracted = opts?.extractArtifacts !== void 0 ? await opts.extractArtifacts(extractionInput) : builtInArtifactDescriptors(activity, ctx.artifactInputs, result);
|
|
486
|
+
if (extracted.length === 0) return [];
|
|
487
|
+
const existingRefs = extracted.filter(isArtifactRef);
|
|
488
|
+
const descriptors = extracted.filter((item) => !isArtifactRef(item));
|
|
489
|
+
if (descriptors.length === 0) return existingRefs;
|
|
490
|
+
if (!persistence.stores.artifacts || !persistence.stores.blobs) throw new Error("Generation artifact persistence requires stores.artifacts and stores.blobs.");
|
|
491
|
+
const refs = [...existingRefs];
|
|
492
|
+
for (const [index, descriptor] of descriptors.entries()) {
|
|
493
|
+
const artifactId = ctx.createId("artifact");
|
|
494
|
+
const resolved = await descriptorBody(descriptor, opts);
|
|
495
|
+
if (!resolved) continue;
|
|
496
|
+
const { body, size, expectedLength, mimeType, sourceUrl } = resolved;
|
|
497
|
+
const name = opts?.nameArtifact?.({
|
|
498
|
+
descriptor: {
|
|
499
|
+
...descriptor,
|
|
500
|
+
mimeType
|
|
501
|
+
},
|
|
502
|
+
activity,
|
|
503
|
+
provider: ctx.provider,
|
|
504
|
+
model: ctx.model,
|
|
505
|
+
threadId,
|
|
506
|
+
runId,
|
|
507
|
+
index
|
|
508
|
+
}) ?? descriptor.name ?? defaultArtifactName({
|
|
509
|
+
...descriptor,
|
|
510
|
+
mimeType
|
|
511
|
+
}, activity, index);
|
|
512
|
+
const key = opts?.storageKey?.({
|
|
513
|
+
artifactId,
|
|
514
|
+
runId,
|
|
515
|
+
threadId,
|
|
516
|
+
role: descriptor.role,
|
|
517
|
+
activity,
|
|
518
|
+
path: descriptor.path,
|
|
519
|
+
mimeType,
|
|
520
|
+
name
|
|
521
|
+
}) ?? artifactBlobKey({
|
|
522
|
+
runId,
|
|
523
|
+
artifactId
|
|
524
|
+
});
|
|
525
|
+
const stored = await persistence.stores.blobs.put(key, body, {
|
|
526
|
+
contentType: mimeType,
|
|
527
|
+
...expectedLength !== void 0 ? { expectedLength } : {},
|
|
528
|
+
customMetadata: {
|
|
529
|
+
runId,
|
|
530
|
+
threadId,
|
|
531
|
+
role: descriptor.role,
|
|
532
|
+
activity,
|
|
533
|
+
path: descriptor.path
|
|
534
|
+
}
|
|
535
|
+
});
|
|
536
|
+
const resolvedSize = size || stored.size || 0;
|
|
537
|
+
const createdAtMs = Date.now();
|
|
538
|
+
const record = {
|
|
539
|
+
artifactId,
|
|
540
|
+
runId,
|
|
541
|
+
threadId,
|
|
542
|
+
blobKey: key,
|
|
543
|
+
name,
|
|
544
|
+
mimeType,
|
|
545
|
+
size: resolvedSize,
|
|
546
|
+
sourceUrl,
|
|
547
|
+
createdAt: createdAtMs
|
|
548
|
+
};
|
|
549
|
+
await persistence.stores.artifacts.save(record);
|
|
550
|
+
refs.push({
|
|
551
|
+
role: descriptor.role,
|
|
552
|
+
artifactId,
|
|
553
|
+
threadId,
|
|
554
|
+
runId,
|
|
555
|
+
name,
|
|
556
|
+
mimeType,
|
|
557
|
+
size: resolvedSize,
|
|
558
|
+
createdAt: new Date(createdAtMs).toISOString(),
|
|
559
|
+
...sourceUrl ? { sourceUrl } : {},
|
|
560
|
+
source: {
|
|
561
|
+
activity,
|
|
562
|
+
path: descriptor.path,
|
|
563
|
+
provider: ctx.provider,
|
|
564
|
+
model: ctx.model,
|
|
565
|
+
mediaType: descriptor.mediaType,
|
|
566
|
+
jobId: descriptor.jobId,
|
|
567
|
+
expiresAt: descriptor.expiresAt instanceof Date ? descriptor.expiresAt.toISOString() : descriptor.expiresAt
|
|
568
|
+
}
|
|
569
|
+
});
|
|
570
|
+
}
|
|
571
|
+
if (opts?.artifactUrl) for (let i = 0; i < refs.length; i++) {
|
|
572
|
+
const ref = refs[i];
|
|
573
|
+
if (ref && !ref.url) {
|
|
574
|
+
const url = opts.artifactUrl(ref);
|
|
575
|
+
if (url) refs[i] = {
|
|
576
|
+
...ref,
|
|
577
|
+
url
|
|
578
|
+
};
|
|
579
|
+
}
|
|
580
|
+
}
|
|
581
|
+
return refs;
|
|
582
|
+
}
|
|
583
|
+
/**
|
|
584
|
+
* Rewrite the live result's media fields to each output ref's durable serve URL
|
|
585
|
+
* (`ref.url`), so the live result matches what a reload restores. Keyed off the
|
|
586
|
+
* ref's `source.path`: `images.<i>` → `result.images[i].url`, `video` →
|
|
587
|
+
* `result.url`, `audio` (object) → `result.audio.url`. tts (a base64 string) and
|
|
588
|
+
* transcription (json) have no media-URL field, so they are left as-is; their
|
|
589
|
+
* durable bytes are reachable via `result.artifacts`. A no-op when no ref has a
|
|
590
|
+
* `url`.
|
|
591
|
+
*/
|
|
592
|
+
function applyDurableMediaUrls(result, refs) {
|
|
593
|
+
let next = result;
|
|
594
|
+
for (const ref of refs) {
|
|
595
|
+
if (ref.role !== "output" || !ref.url) continue;
|
|
596
|
+
const path = ref.source.path;
|
|
597
|
+
if (path.startsWith("images.")) {
|
|
598
|
+
const index = Number(path.slice(7));
|
|
599
|
+
const images = next.images;
|
|
600
|
+
if (Array.isArray(images) && objectValue(images[index])) {
|
|
601
|
+
const cloned = [...images];
|
|
602
|
+
cloned[index] = {
|
|
603
|
+
...objectValue(images[index]),
|
|
604
|
+
url: ref.url
|
|
605
|
+
};
|
|
606
|
+
next = {
|
|
607
|
+
...next,
|
|
608
|
+
images: cloned
|
|
609
|
+
};
|
|
610
|
+
}
|
|
611
|
+
} else if (path === "video") next = {
|
|
612
|
+
...next,
|
|
613
|
+
url: ref.url
|
|
614
|
+
};
|
|
615
|
+
else if (path === "audio" && objectValue(next.audio)) next = {
|
|
616
|
+
...next,
|
|
617
|
+
audio: {
|
|
618
|
+
...objectValue(next.audio),
|
|
619
|
+
url: ref.url
|
|
620
|
+
}
|
|
621
|
+
};
|
|
622
|
+
}
|
|
623
|
+
return next;
|
|
624
|
+
}
|
|
625
|
+
function resolvePersistencePlan(persistence) {
|
|
626
|
+
return {
|
|
627
|
+
wantsInterrupts: persistence.stores.interrupts !== void 0,
|
|
628
|
+
wantsArtifactPersistence: persistence.stores.artifacts !== void 0 && persistence.stores.blobs !== void 0,
|
|
629
|
+
runs: persistence.stores.runs
|
|
630
|
+
};
|
|
631
|
+
}
|
|
632
|
+
async function createOrResumeRun(runs, runId, threadId) {
|
|
633
|
+
await runs?.createOrResume({
|
|
634
|
+
runId,
|
|
635
|
+
threadId,
|
|
636
|
+
startedAt: Date.now()
|
|
637
|
+
});
|
|
638
|
+
}
|
|
639
|
+
async function completeRun(runs, runId, usage) {
|
|
640
|
+
await runs?.update(runId, {
|
|
641
|
+
status: "completed",
|
|
642
|
+
finishedAt: Date.now(),
|
|
643
|
+
...usage ? { usage } : {}
|
|
644
|
+
});
|
|
645
|
+
}
|
|
646
|
+
async function failRun(runs, runId, error) {
|
|
647
|
+
await runs?.update(runId, {
|
|
648
|
+
status: "failed",
|
|
649
|
+
finishedAt: Date.now(),
|
|
650
|
+
error: { message: error instanceof Error ? error.message : String(error) }
|
|
651
|
+
});
|
|
652
|
+
}
|
|
653
|
+
/**
|
|
654
|
+
* Record a human-in-the-loop PAUSE.
|
|
655
|
+
*
|
|
656
|
+
* Deliberately writes NO `finishedAt`: `'interrupted'` is not a terminal status
|
|
657
|
+
* (`isTerminalRunStatus('interrupted')` is `false`), and stamping a terminal
|
|
658
|
+
* timestamp on it told every reader the run was over while it was in fact
|
|
659
|
+
* waiting for a human. Only `abortRun`/`completeRun`/`failRun` finish a run.
|
|
660
|
+
*/
|
|
661
|
+
async function interruptRun(runs, runId) {
|
|
662
|
+
await runs?.update(runId, { status: "interrupted" });
|
|
663
|
+
}
|
|
664
|
+
/**
|
|
665
|
+
* Record that the run has ended for good — an explicit cancel, or a disconnect
|
|
666
|
+
* on a run that has nothing to reattach to. Terminal, so it carries
|
|
667
|
+
* `finishedAt`.
|
|
668
|
+
*/
|
|
669
|
+
async function abortRun(runs, runId) {
|
|
670
|
+
await runs?.update(runId, {
|
|
671
|
+
status: "aborted",
|
|
672
|
+
finishedAt: Date.now()
|
|
673
|
+
});
|
|
674
|
+
}
|
|
675
|
+
/**
|
|
676
|
+
* Whether some middleware has declared this run detachable — i.e. it has a
|
|
677
|
+
* durable event log and a run store, so a disconnect can be survived and the
|
|
678
|
+
* run picked back up rather than destroyed.
|
|
679
|
+
*
|
|
680
|
+
* The capability is read from CORE, never from `@tanstack/ai-sandbox`: sandbox
|
|
681
|
+
* provides it, persistence consumes it, and a persistence → sandbox import
|
|
682
|
+
* would invert the layering.
|
|
683
|
+
*/
|
|
684
|
+
function detachableRun(ctx) {
|
|
685
|
+
return getDetachableRun(ctx, { optional: true }) === true;
|
|
686
|
+
}
|
|
687
|
+
/**
|
|
688
|
+
* @param persistence - Must satisfy {@link ChatTranscriptStores} (messages
|
|
689
|
+
* required). Known-absent `messages` or `interrupts` without `runs` fail at
|
|
690
|
+
* compile time; fully dynamic bags are checked at runtime.
|
|
691
|
+
*/
|
|
692
|
+
function withPersistence(persistence, options = {}) {
|
|
693
|
+
validateChatPersistenceStores(persistence);
|
|
694
|
+
const snapshotStreaming = options.snapshotStreaming ?? false;
|
|
695
|
+
const snapshotIntervalMs = options.snapshotIntervalMs ?? 1e3;
|
|
696
|
+
const { wantsInterrupts, runs } = resolvePersistencePlan(persistence);
|
|
697
|
+
const messageStore = persistence.stores.messages;
|
|
698
|
+
if (!messageStore) throw new Error("Chat persistence requires stores.messages.");
|
|
699
|
+
return defineChatMiddleware({
|
|
700
|
+
name: "chat-persistence",
|
|
701
|
+
provides: [PersistenceCapability, ...wantsInterrupts ? [InterruptsCapability] : []],
|
|
702
|
+
setup(ctx) {
|
|
703
|
+
providePersistence(ctx, persistence);
|
|
704
|
+
runState.set(ctx, {
|
|
705
|
+
merged: false,
|
|
706
|
+
interrupted: false
|
|
707
|
+
});
|
|
708
|
+
if (wantsInterrupts && persistence.stores.interrupts) provideInterrupts(ctx, persistence.stores.interrupts);
|
|
709
|
+
providePendingTurn(ctx, { snapshot: async () => {
|
|
710
|
+
const stored = await messageStore.loadThread(ctx.threadId);
|
|
711
|
+
const list = ctx.messages.length > 0 ? [...ctx.messages] : stored;
|
|
712
|
+
await messageStore.saveThread(ctx.threadId, list);
|
|
713
|
+
} });
|
|
714
|
+
},
|
|
715
|
+
async onConfig(ctx, config) {
|
|
716
|
+
if (ctx.phase !== "init") return;
|
|
717
|
+
const patch = {};
|
|
718
|
+
if (wantsInterrupts && persistence.stores.interrupts) {
|
|
719
|
+
const pending = await persistence.stores.interrupts.listPending(ctx.threadId);
|
|
720
|
+
const resumeByInterruptId = validatePendingResumes(pending, config.resume);
|
|
721
|
+
if ((config.resume?.length ?? 0) > 0) {
|
|
722
|
+
const resumeToolState = resumeToolStateFromPending(pending, resumeByInterruptId);
|
|
723
|
+
patch.resume = [];
|
|
724
|
+
if (resumeToolState) patch.resumeToolState = resumeToolState;
|
|
725
|
+
}
|
|
726
|
+
const state = runState.get(ctx);
|
|
727
|
+
if (state && pending.length > 0) state.pendingResumes = {
|
|
728
|
+
pending,
|
|
729
|
+
resumeByInterruptId
|
|
730
|
+
};
|
|
731
|
+
}
|
|
732
|
+
await createOrResumeRun(runs, ctx.runId, ctx.threadId);
|
|
733
|
+
{
|
|
734
|
+
const state = runState.get(ctx);
|
|
735
|
+
if (!state?.merged) {
|
|
736
|
+
if (state) state.merged = true;
|
|
737
|
+
const stored = await messageStore.loadThread(ctx.threadId);
|
|
738
|
+
patch.messages = config.messages.length > 0 ? config.messages : stored;
|
|
739
|
+
}
|
|
740
|
+
}
|
|
741
|
+
return Object.keys(patch).length > 0 ? patch : void 0;
|
|
742
|
+
},
|
|
743
|
+
async onStart(ctx) {
|
|
744
|
+
try {
|
|
745
|
+
await messageStore.saveThread(ctx.threadId, [...ctx.messages]);
|
|
746
|
+
} catch {}
|
|
747
|
+
},
|
|
748
|
+
async onChunk(ctx, chunk) {
|
|
749
|
+
if (chunk.type === "TEXT_MESSAGE_START") {
|
|
750
|
+
const s = runState.get(ctx);
|
|
751
|
+
if (s) {
|
|
752
|
+
s.streamingMessageId = chunk.messageId;
|
|
753
|
+
s.streamingText = "";
|
|
754
|
+
}
|
|
755
|
+
}
|
|
756
|
+
if (snapshotStreaming && chunk.type === "TEXT_MESSAGE_CONTENT" && typeof chunk.delta === "string") {
|
|
757
|
+
const snapshotState = runState.get(ctx);
|
|
758
|
+
if (snapshotState) {
|
|
759
|
+
snapshotState.streamingText = (snapshotState.streamingText ?? "") + chunk.delta;
|
|
760
|
+
const now = Date.now();
|
|
761
|
+
if (now - (snapshotState.lastSnapshotAt ?? 0) >= snapshotIntervalMs) {
|
|
762
|
+
snapshotState.lastSnapshotAt = now;
|
|
763
|
+
try {
|
|
764
|
+
await messageStore.saveThread(ctx.threadId, [...ctx.messages, {
|
|
765
|
+
role: "assistant",
|
|
766
|
+
content: snapshotState.streamingText,
|
|
767
|
+
...snapshotState.streamingMessageId ? { id: snapshotState.streamingMessageId } : {}
|
|
768
|
+
}]);
|
|
769
|
+
} catch {}
|
|
770
|
+
}
|
|
771
|
+
}
|
|
772
|
+
}
|
|
773
|
+
if (chunk.type !== "RUN_FINISHED" || chunk.outcome?.type !== "interrupt") return;
|
|
774
|
+
const state = runState.get(ctx);
|
|
775
|
+
if (!state) return;
|
|
776
|
+
if (wantsInterrupts && persistence.stores.interrupts) {
|
|
777
|
+
await commitPendingResumes(state, persistence.stores.interrupts);
|
|
778
|
+
for (const interrupt of chunk.outcome.interrupts) await persistence.stores.interrupts.create({
|
|
779
|
+
interruptId: interrupt.id,
|
|
780
|
+
runId: ctx.runId,
|
|
781
|
+
threadId: ctx.threadId,
|
|
782
|
+
requestedAt: Date.now(),
|
|
783
|
+
payload: interruptPayload(interrupt)
|
|
784
|
+
});
|
|
785
|
+
}
|
|
786
|
+
await interruptRun(runs, ctx.runId);
|
|
787
|
+
await messageStore.saveThread(ctx.threadId, [...ctx.messages]);
|
|
788
|
+
state.interrupted = true;
|
|
789
|
+
},
|
|
790
|
+
async onFinish(ctx, info) {
|
|
791
|
+
const state = runState.get(ctx);
|
|
792
|
+
if (state?.interrupted) return;
|
|
793
|
+
await messageStore.saveThread(ctx.threadId, finishedTranscript(ctx.messages, info, state?.streamingMessageId));
|
|
794
|
+
await completeRun(runs, ctx.runId, info.usage);
|
|
795
|
+
await commitPendingResumes(state, persistence.stores.interrupts);
|
|
796
|
+
},
|
|
797
|
+
async onError(ctx, info) {
|
|
798
|
+
await failRun(runs, ctx.runId, info.error);
|
|
799
|
+
},
|
|
800
|
+
async onAbort(ctx, info) {
|
|
801
|
+
const cancelled = info.cancelRequested === true || runs !== void 0 && await wasCancelRequested(runs, ctx.runId);
|
|
802
|
+
const state = runState.get(ctx);
|
|
803
|
+
if (cancelled || !detachableRun(ctx) && state?.interrupted !== true) {
|
|
804
|
+
await abortRun(runs, ctx.runId);
|
|
805
|
+
return;
|
|
806
|
+
}
|
|
807
|
+
}
|
|
808
|
+
});
|
|
809
|
+
}
|
|
810
|
+
function withGenerationPersistence(persistence, opts = {}) {
|
|
811
|
+
validateGenerationPersistenceStores(persistence);
|
|
812
|
+
const { wantsArtifactPersistence } = resolvePersistencePlan(persistence);
|
|
813
|
+
const generationRuns = persistence.stores.generationRuns;
|
|
814
|
+
if (!generationRuns) throw new Error("Generation persistence requires stores.generationRuns.");
|
|
815
|
+
const runIdOf = (ctx) => ctx.runId ?? ctx.requestId;
|
|
816
|
+
return {
|
|
817
|
+
name: "generation-persistence",
|
|
818
|
+
async onStart(ctx) {
|
|
819
|
+
const runId = runIdOf(ctx);
|
|
820
|
+
await generationRuns.createOrResume({
|
|
821
|
+
runId,
|
|
822
|
+
activity: ctx.activity,
|
|
823
|
+
provider: ctx.provider,
|
|
824
|
+
model: ctx.model,
|
|
825
|
+
startedAt: Date.now(),
|
|
826
|
+
threadId: generationScope(ctx, opts)
|
|
827
|
+
});
|
|
828
|
+
if (wantsArtifactPersistence) ctx.resultTransforms?.push(async (result) => {
|
|
829
|
+
const refs = await persistGenerationArtifacts(persistence, opts, ctx, result);
|
|
830
|
+
if (refs.length === 0) return void 0;
|
|
831
|
+
const base = objectValue(result) ?? {};
|
|
832
|
+
const existing = base.artifacts;
|
|
833
|
+
return applyDurableMediaUrls({
|
|
834
|
+
...base,
|
|
835
|
+
artifacts: [...Array.isArray(existing) ? existing : [], ...refs]
|
|
836
|
+
}, refs);
|
|
837
|
+
});
|
|
838
|
+
ctx.resultTransforms?.push(async (result) => {
|
|
839
|
+
const rawArtifacts = objectValue(result)?.artifacts;
|
|
840
|
+
const artifacts = Array.isArray(rawArtifacts) ? rawArtifacts.filter(isArtifactRef) : [];
|
|
841
|
+
await generationRuns.update(runId, {
|
|
842
|
+
result,
|
|
843
|
+
...artifacts.length > 0 ? { artifacts } : {}
|
|
844
|
+
});
|
|
845
|
+
});
|
|
846
|
+
},
|
|
847
|
+
async onFinish(ctx, info) {
|
|
848
|
+
await generationRuns.update(runIdOf(ctx), {
|
|
849
|
+
status: "completed",
|
|
850
|
+
finishedAt: Date.now(),
|
|
851
|
+
...info.usage ? { usage: info.usage } : {}
|
|
852
|
+
});
|
|
853
|
+
},
|
|
854
|
+
async onError(ctx, info) {
|
|
855
|
+
await generationRuns.update(runIdOf(ctx), {
|
|
856
|
+
status: "failed",
|
|
857
|
+
finishedAt: Date.now(),
|
|
858
|
+
error: { message: info.error instanceof Error ? info.error.message : String(info.error) }
|
|
859
|
+
});
|
|
860
|
+
},
|
|
861
|
+
async onAbort(ctx, _info) {
|
|
862
|
+
await generationRuns.update(runIdOf(ctx), {
|
|
863
|
+
status: "aborted",
|
|
864
|
+
finishedAt: Date.now()
|
|
865
|
+
});
|
|
866
|
+
}
|
|
867
|
+
};
|
|
868
|
+
}
|
|
869
|
+
//#endregion
|
|
870
|
+
export { abortRun, interruptRun, withGenerationPersistence, withPersistence };
|
|
871
|
+
|
|
872
|
+
//# sourceMappingURL=middleware.js.map
|