@convex-dev/agent 0.0.1-alpha.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.
Files changed (129) hide show
  1. package/LICENSE +201 -0
  2. package/README.md +55 -0
  3. package/dist/commonjs/client/index.d.ts +198 -0
  4. package/dist/commonjs/client/index.d.ts.map +1 -0
  5. package/dist/commonjs/client/index.js +365 -0
  6. package/dist/commonjs/client/index.js.map +1 -0
  7. package/dist/commonjs/client/types.d.ts +21 -0
  8. package/dist/commonjs/client/types.d.ts.map +1 -0
  9. package/dist/commonjs/client/types.js +2 -0
  10. package/dist/commonjs/client/types.js.map +1 -0
  11. package/dist/commonjs/component/_generated/api.d.ts +12 -0
  12. package/dist/commonjs/component/_generated/api.d.ts.map +1 -0
  13. package/dist/commonjs/component/_generated/api.js +22 -0
  14. package/dist/commonjs/component/_generated/api.js.map +1 -0
  15. package/dist/commonjs/component/_generated/server.d.ts +64 -0
  16. package/dist/commonjs/component/_generated/server.d.ts.map +1 -0
  17. package/dist/commonjs/component/_generated/server.js +74 -0
  18. package/dist/commonjs/component/_generated/server.js.map +1 -0
  19. package/dist/commonjs/component/convex.config.d.ts +3 -0
  20. package/dist/commonjs/component/convex.config.d.ts.map +1 -0
  21. package/dist/commonjs/component/convex.config.js +3 -0
  22. package/dist/commonjs/component/convex.config.js.map +1 -0
  23. package/dist/commonjs/component/lib.d.ts +2 -0
  24. package/dist/commonjs/component/lib.d.ts.map +1 -0
  25. package/dist/commonjs/component/lib.js +2 -0
  26. package/dist/commonjs/component/lib.js.map +1 -0
  27. package/dist/commonjs/component/messages.d.ts +1913 -0
  28. package/dist/commonjs/component/messages.d.ts.map +1 -0
  29. package/dist/commonjs/component/messages.js +787 -0
  30. package/dist/commonjs/component/messages.js.map +1 -0
  31. package/dist/commonjs/component/schema.d.ts +5496 -0
  32. package/dist/commonjs/component/schema.d.ts.map +1 -0
  33. package/dist/commonjs/component/schema.js +97 -0
  34. package/dist/commonjs/component/schema.js.map +1 -0
  35. package/dist/commonjs/component/vector/tables.d.ts +40 -0
  36. package/dist/commonjs/component/vector/tables.d.ts.map +1 -0
  37. package/dist/commonjs/component/vector/tables.js +46 -0
  38. package/dist/commonjs/component/vector/tables.js.map +1 -0
  39. package/dist/commonjs/mapping.d.ts +26 -0
  40. package/dist/commonjs/mapping.d.ts.map +1 -0
  41. package/dist/commonjs/mapping.js +101 -0
  42. package/dist/commonjs/mapping.js.map +1 -0
  43. package/dist/commonjs/package.json +3 -0
  44. package/dist/commonjs/react/index.d.ts +2 -0
  45. package/dist/commonjs/react/index.d.ts.map +1 -0
  46. package/dist/commonjs/react/index.js +8 -0
  47. package/dist/commonjs/react/index.js.map +1 -0
  48. package/dist/commonjs/shared.d.ts +9 -0
  49. package/dist/commonjs/shared.d.ts.map +1 -0
  50. package/dist/commonjs/shared.js +29 -0
  51. package/dist/commonjs/shared.js.map +1 -0
  52. package/dist/commonjs/validators.d.ts +6177 -0
  53. package/dist/commonjs/validators.d.ts.map +1 -0
  54. package/dist/commonjs/validators.js +171 -0
  55. package/dist/commonjs/validators.js.map +1 -0
  56. package/dist/esm/client/index.d.ts +198 -0
  57. package/dist/esm/client/index.d.ts.map +1 -0
  58. package/dist/esm/client/index.js +365 -0
  59. package/dist/esm/client/index.js.map +1 -0
  60. package/dist/esm/client/types.d.ts +21 -0
  61. package/dist/esm/client/types.d.ts.map +1 -0
  62. package/dist/esm/client/types.js +2 -0
  63. package/dist/esm/client/types.js.map +1 -0
  64. package/dist/esm/component/_generated/api.d.ts +12 -0
  65. package/dist/esm/component/_generated/api.d.ts.map +1 -0
  66. package/dist/esm/component/_generated/api.js +22 -0
  67. package/dist/esm/component/_generated/api.js.map +1 -0
  68. package/dist/esm/component/_generated/server.d.ts +64 -0
  69. package/dist/esm/component/_generated/server.d.ts.map +1 -0
  70. package/dist/esm/component/_generated/server.js +74 -0
  71. package/dist/esm/component/_generated/server.js.map +1 -0
  72. package/dist/esm/component/convex.config.d.ts +3 -0
  73. package/dist/esm/component/convex.config.d.ts.map +1 -0
  74. package/dist/esm/component/convex.config.js +3 -0
  75. package/dist/esm/component/convex.config.js.map +1 -0
  76. package/dist/esm/component/lib.d.ts +2 -0
  77. package/dist/esm/component/lib.d.ts.map +1 -0
  78. package/dist/esm/component/lib.js +2 -0
  79. package/dist/esm/component/lib.js.map +1 -0
  80. package/dist/esm/component/messages.d.ts +1913 -0
  81. package/dist/esm/component/messages.d.ts.map +1 -0
  82. package/dist/esm/component/messages.js +787 -0
  83. package/dist/esm/component/messages.js.map +1 -0
  84. package/dist/esm/component/schema.d.ts +5496 -0
  85. package/dist/esm/component/schema.d.ts.map +1 -0
  86. package/dist/esm/component/schema.js +97 -0
  87. package/dist/esm/component/schema.js.map +1 -0
  88. package/dist/esm/component/vector/tables.d.ts +40 -0
  89. package/dist/esm/component/vector/tables.d.ts.map +1 -0
  90. package/dist/esm/component/vector/tables.js +46 -0
  91. package/dist/esm/component/vector/tables.js.map +1 -0
  92. package/dist/esm/mapping.d.ts +26 -0
  93. package/dist/esm/mapping.d.ts.map +1 -0
  94. package/dist/esm/mapping.js +101 -0
  95. package/dist/esm/mapping.js.map +1 -0
  96. package/dist/esm/package.json +3 -0
  97. package/dist/esm/react/index.d.ts +2 -0
  98. package/dist/esm/react/index.d.ts.map +1 -0
  99. package/dist/esm/react/index.js +8 -0
  100. package/dist/esm/react/index.js.map +1 -0
  101. package/dist/esm/shared.d.ts +9 -0
  102. package/dist/esm/shared.d.ts.map +1 -0
  103. package/dist/esm/shared.js +29 -0
  104. package/dist/esm/shared.js.map +1 -0
  105. package/dist/esm/validators.d.ts +6177 -0
  106. package/dist/esm/validators.d.ts.map +1 -0
  107. package/dist/esm/validators.js +171 -0
  108. package/dist/esm/validators.js.map +1 -0
  109. package/package.json +91 -0
  110. package/react/package.json +5 -0
  111. package/src/client/index.ts +659 -0
  112. package/src/client/types.ts +54 -0
  113. package/src/component/_generated/api.d.ts +1497 -0
  114. package/src/component/_generated/api.js +23 -0
  115. package/src/component/_generated/dataModel.d.ts +60 -0
  116. package/src/component/_generated/server.d.ts +149 -0
  117. package/src/component/_generated/server.js +90 -0
  118. package/src/component/convex.config.ts +3 -0
  119. package/src/component/lib.test.ts +13 -0
  120. package/src/component/lib.ts +2 -0
  121. package/src/component/messages.ts +959 -0
  122. package/src/component/schema.ts +101 -0
  123. package/src/component/setup.test.ts +5 -0
  124. package/src/component/vector/tables.ts +92 -0
  125. package/src/mapping.ts +160 -0
  126. package/src/react/index.ts +8 -0
  127. package/src/shared.ts +35 -0
  128. package/src/validators.test.ts +101 -0
  129. package/src/validators.ts +258 -0
@@ -0,0 +1,787 @@
1
+ import { assert, omit, pick } from "convex-helpers";
2
+ import { paginator } from "convex-helpers/server/pagination";
3
+ import { mergedStream, stream } from "convex-helpers/server/stream";
4
+ import { nullable, partial } from "convex-helpers/validators";
5
+ import { DEFAULT_MESSAGE_RANGE, extractText, isTool } from "../shared.js";
6
+ import { vChatStatus, vMessageStatus, vMessageWithFileAndId, vSearchOptions, vStepWithMessagesWithFileAndId, } from "../validators.js";
7
+ import { api, internal } from "./_generated/api.js";
8
+ import { action, internalMutation, internalQuery, mutation, query, } from "./_generated/server.js";
9
+ import { schema, v } from "./schema.js";
10
+ import { getVectorTableName, VectorDimensions, vVectorId, } from "./vector/tables.js";
11
+ export const getChat = query({
12
+ args: { chatId: v.id("chats") },
13
+ handler: async (ctx, args) => {
14
+ return ctx.db.get(args.chatId);
15
+ },
16
+ returns: v.union(v.doc("chats"), v.null()),
17
+ });
18
+ export const getChatsByUserId = query({
19
+ args: {
20
+ userId: v.string(),
21
+ // Note: the other arguments cannot change from when the cursor was created.
22
+ cursor: v.optional(v.union(v.string(), v.null())),
23
+ limit: v.optional(v.number()),
24
+ offset: v.optional(v.number()),
25
+ statuses: v.optional(v.array(vChatStatus)),
26
+ },
27
+ handler: async (ctx, args) => {
28
+ const streams = (args.statuses ?? ["active"]).map((status) => stream(ctx.db, schema)
29
+ .query("chats")
30
+ .withIndex("status_userId_order", (q) => q
31
+ .eq("status", status)
32
+ .eq("userId", args.userId)
33
+ .gte("order", args.offset ?? 0)));
34
+ const chats = await mergedStream(streams, ["order", "stepOrder"]).paginate({
35
+ numItems: args.limit ?? 100,
36
+ cursor: args.cursor ?? null,
37
+ });
38
+ return {
39
+ chats: chats.page,
40
+ continueCursor: chats.continueCursor,
41
+ isDone: chats.isDone,
42
+ };
43
+ },
44
+ returns: v.object({
45
+ chats: v.array(v.doc("chats")),
46
+ continueCursor: v.string(),
47
+ isDone: v.boolean(),
48
+ }),
49
+ });
50
+ const vChat = schema.tables.chats.validator;
51
+ const statuses = vChat.fields.status.members.map((m) => m.value);
52
+ export const createChat = mutation({
53
+ args: omit(vChat.fields, ["order", "status"]),
54
+ handler: async (ctx, args) => {
55
+ const streams = statuses.map((status) => stream(ctx.db, schema)
56
+ .query("chats")
57
+ .withIndex("status_userId_order", (q) => q.eq("status", status).eq("userId", args.userId))
58
+ .order("desc"));
59
+ const latestChat = await mergedStream(streams, ["order"]).first();
60
+ const order = (latestChat?.order ?? -1) + 1;
61
+ const chatId = await ctx.db.insert("chats", {
62
+ ...args,
63
+ order,
64
+ status: "active",
65
+ });
66
+ return (await ctx.db.get(chatId));
67
+ },
68
+ returns: v.doc("chats"),
69
+ });
70
+ export const updateChat = mutation({
71
+ args: {
72
+ chatId: v.id("chats"),
73
+ patch: v.object(partial(pick(vChat.fields, [
74
+ "title",
75
+ "summary",
76
+ "defaultSystemPrompt",
77
+ "status",
78
+ ]))),
79
+ },
80
+ handler: async (ctx, args) => {
81
+ const chat = await ctx.db.get(args.chatId);
82
+ assert(chat, `Chat ${args.chatId} not found`);
83
+ await ctx.db.patch(args.chatId, args.patch);
84
+ return (await ctx.db.get(args.chatId));
85
+ },
86
+ returns: v.doc("chats"),
87
+ });
88
+ export const archiveChat = mutation({
89
+ args: { chatId: v.id("chats") },
90
+ handler: async (ctx, args) => {
91
+ const chat = await ctx.db.get(args.chatId);
92
+ assert(chat, `Chat ${args.chatId} not found`);
93
+ await ctx.db.patch(args.chatId, { status: "archived" });
94
+ return (await ctx.db.get(args.chatId));
95
+ },
96
+ returns: v.doc("chats"),
97
+ });
98
+ export const deleteAllForUserId = action({
99
+ args: { userId: v.string() },
100
+ handler: async (ctx, args) => {
101
+ let messagesCursor = null;
102
+ let chatsCursor = null;
103
+ let isDone = false;
104
+ while (!isDone) {
105
+ const result = await ctx.runMutation(internal.messages._deletePageForUserId, {
106
+ userId: args.userId,
107
+ messagesCursor,
108
+ chatsCursor,
109
+ });
110
+ messagesCursor = result.messagesCursor;
111
+ chatsCursor = result.chatsCursor;
112
+ isDone = result.isDone;
113
+ }
114
+ },
115
+ returns: v.null(),
116
+ });
117
+ export const deleteAllForUserIdAsync = mutation({
118
+ args: {
119
+ userId: v.string(),
120
+ },
121
+ handler: async (ctx, args) => {
122
+ const isDone = await deleteAllFroUserIdAsyncHandler(ctx, {
123
+ userId: args.userId,
124
+ messagesCursor: null,
125
+ chatsCursor: null,
126
+ });
127
+ return isDone;
128
+ },
129
+ returns: v.boolean(),
130
+ });
131
+ const deleteAllArgs = {
132
+ userId: v.string(),
133
+ messagesCursor: nullable(v.string()),
134
+ chatsCursor: nullable(v.string()),
135
+ };
136
+ const deleteAllReturns = {
137
+ messagesCursor: v.string(),
138
+ chatsCursor: nullable(v.string()),
139
+ isDone: v.boolean(),
140
+ };
141
+ export const _deleteAllForUserIdAsync = internalMutation({
142
+ args: deleteAllArgs,
143
+ handler: deleteAllFroUserIdAsyncHandler,
144
+ returns: v.boolean(),
145
+ });
146
+ async function deleteAllFroUserIdAsyncHandler(ctx, args) {
147
+ const result = await deletePageForUserId(ctx, args);
148
+ if (!result.isDone) {
149
+ await ctx.scheduler.runAfter(0, internal.messages._deleteAllForUserIdAsync, {
150
+ userId: args.userId,
151
+ messagesCursor: result.messagesCursor,
152
+ chatsCursor: result.chatsCursor,
153
+ });
154
+ }
155
+ return result.isDone;
156
+ }
157
+ export const _deletePageForUserId = internalMutation({
158
+ args: deleteAllArgs,
159
+ handler: deletePageForUserId,
160
+ returns: deleteAllReturns,
161
+ });
162
+ async function deletePageForUserId(ctx, args) {
163
+ const streams = statuses.map((status) => stream(ctx.db, schema)
164
+ .query("chats")
165
+ .withIndex("status_userId_order", (q) => q.eq("status", status).eq("userId", args.userId))
166
+ .order("desc"));
167
+ const chatStreams = mergedStream(streams, ["order"]);
168
+ const messages = await chatStreams
169
+ .flatMap(async (c) => stream(ctx.db, schema)
170
+ .query("messages")
171
+ .withIndex("chatId_status_tool_order_stepOrder", (q) => q.eq("chatId", c._id).eq("status", "success")), ["tool", "order", "stepOrder"])
172
+ .paginate({
173
+ numItems: 100,
174
+ cursor: args.messagesCursor ?? null,
175
+ });
176
+ await Promise.all(messages.page.map((m) => deleteMessage(ctx, m)));
177
+ if (messages.isDone) {
178
+ const chats = await chatStreams.paginate({
179
+ numItems: 100,
180
+ cursor: args.chatsCursor ?? null,
181
+ });
182
+ await Promise.all(chats.page.map((c) => ctx.db.delete(c._id)));
183
+ return {
184
+ messagesCursor: messages.continueCursor,
185
+ chatsCursor: chats.continueCursor,
186
+ isDone: chats.isDone,
187
+ };
188
+ }
189
+ return {
190
+ messagesCursor: messages.continueCursor,
191
+ chatsCursor: null,
192
+ isDone: messages.isDone,
193
+ };
194
+ }
195
+ async function deleteMessage(ctx, messageDoc) {
196
+ await ctx.db.delete(messageDoc._id);
197
+ if (messageDoc.fileId) {
198
+ const file = await ctx.db.get(messageDoc.fileId);
199
+ if (file) {
200
+ await ctx.db.patch(messageDoc.fileId, { refcount: file.refcount - 1 });
201
+ }
202
+ }
203
+ }
204
+ const deleteChatArgs = {
205
+ chatId: v.id("chats"),
206
+ cursor: v.optional(v.string()),
207
+ limit: v.optional(v.number()),
208
+ };
209
+ const deleteChatReturns = {
210
+ cursor: v.string(),
211
+ isDone: v.boolean(),
212
+ };
213
+ export const deleteAllForChatIdSync = action({
214
+ args: deleteChatArgs,
215
+ handler: async (ctx, args) => {
216
+ const result = await ctx.runMutation(internal.messages._deletePageForChatId, { chatId: args.chatId, cursor: args.cursor, limit: args.limit });
217
+ return result;
218
+ },
219
+ returns: deleteChatReturns,
220
+ });
221
+ export const deleteAllForChatIdAsync = mutation({
222
+ args: deleteChatArgs,
223
+ handler: async (ctx, args) => {
224
+ const result = await deletePageForChatIdHandler(ctx, args);
225
+ if (!result.isDone) {
226
+ await ctx.scheduler.runAfter(0, api.messages.deleteAllForChatIdAsync, {
227
+ chatId: args.chatId,
228
+ cursor: result.cursor,
229
+ });
230
+ }
231
+ return result;
232
+ },
233
+ returns: deleteChatReturns,
234
+ });
235
+ export const _deletePageForChatId = internalMutation({
236
+ args: deleteChatArgs,
237
+ handler: deletePageForChatIdHandler,
238
+ returns: deleteChatReturns,
239
+ });
240
+ async function deletePageForChatIdHandler(ctx, args) {
241
+ const messages = await stream(ctx.db, schema)
242
+ .query("messages")
243
+ .withIndex("chatId_status_tool_order_stepOrder", (q) => q.eq("chatId", args.chatId).eq("status", "success"))
244
+ .paginate({
245
+ numItems: args.limit ?? 100,
246
+ cursor: args.cursor ?? null,
247
+ });
248
+ await Promise.all(messages.page.map((m) => deleteMessage(ctx, m)));
249
+ await ctx.db.delete(args.chatId);
250
+ return {
251
+ cursor: messages.continueCursor,
252
+ isDone: messages.isDone,
253
+ };
254
+ }
255
+ export const getFilesToDelete = query({
256
+ args: {
257
+ cursor: v.optional(v.string()),
258
+ limit: v.optional(v.number()),
259
+ },
260
+ handler: async (ctx, args) => {
261
+ const files = await paginator(ctx.db, schema)
262
+ .query("files")
263
+ .withIndex("refcount", (q) => q.eq("refcount", 0))
264
+ .paginate({
265
+ numItems: args.limit ?? 100,
266
+ cursor: args.cursor ?? null,
267
+ });
268
+ return {
269
+ files: files.page,
270
+ continueCursor: files.continueCursor,
271
+ isDone: files.isDone,
272
+ };
273
+ },
274
+ returns: v.object({
275
+ files: v.array(v.doc("files")),
276
+ continueCursor: v.string(),
277
+ isDone: v.boolean(),
278
+ }),
279
+ });
280
+ export const vMessageDoc = schema.tables.messages.validator;
281
+ export const messageStatuses = vMessageDoc.fields.status.members.map((m) => m.value);
282
+ const addMessagesArgs = {
283
+ chatId: v.id("chats"),
284
+ stepId: v.optional(v.id("steps")),
285
+ parentMessageId: v.optional(v.id("messages")),
286
+ messages: v.array(vMessageWithFileAndId),
287
+ model: v.optional(v.string()),
288
+ agentName: v.optional(v.string()),
289
+ pending: v.optional(v.boolean()),
290
+ failPendingSteps: v.optional(v.boolean()),
291
+ };
292
+ export const addMessages = mutation({
293
+ args: addMessagesArgs,
294
+ handler: addMessagesHandler,
295
+ returns: v.object({
296
+ messages: v.array(v.doc("messages")),
297
+ pending: v.optional(v.doc("messages")),
298
+ }),
299
+ });
300
+ async function addMessagesHandler(ctx, args) {
301
+ const chat = await ctx.db.get(args.chatId);
302
+ assert(chat, `Chat ${args.chatId} not found`);
303
+ const { failPendingSteps, pending, messages, parentMessageId, ...rest } = args;
304
+ if (failPendingSteps) {
305
+ const pendingMessages = await ctx.db
306
+ .query("messages")
307
+ .withIndex("chatId_status_tool_order_stepOrder", (q) => q.eq("chatId", args.chatId).eq("status", "pending"))
308
+ .collect();
309
+ await Promise.all(pendingMessages.map((m) => ctx.db.patch(m._id, { status: "failed", text: "Restarting" })));
310
+ }
311
+ let order;
312
+ const maxMessage = await getMaxMessage(ctx, args.chatId);
313
+ // If the previous message isn't our parent, we make a new thread.
314
+ const threadId = parentMessageId && maxMessage?._id === parentMessageId
315
+ ? maxMessage.threadId ?? parentMessageId
316
+ : parentMessageId;
317
+ order = maxMessage?.order ?? -1;
318
+ const toReturn = [];
319
+ if (messages.length > 0) {
320
+ for (const { message, fileId, id } of messages) {
321
+ const tool = isTool(message);
322
+ if (!tool) {
323
+ order++;
324
+ }
325
+ const text = extractText(message);
326
+ const messageId = await ctx.db.insert("messages", {
327
+ ...rest,
328
+ threadId,
329
+ userId: chat.userId,
330
+ message,
331
+ id,
332
+ order,
333
+ tool,
334
+ text,
335
+ fileId,
336
+ status: pending ? "pending" : "success",
337
+ });
338
+ toReturn.push((await ctx.db.get(messageId)));
339
+ }
340
+ }
341
+ return { messages: toReturn };
342
+ }
343
+ async function getMaxMessage(ctx, chatId) {
344
+ return mergedStream(["success", "pending"].map((status) => stream(ctx.db, schema)
345
+ .query("messages")
346
+ .withIndex("chatId_status_tool_order_stepOrder", (q) => q.eq("chatId", chatId).eq("status", status).eq("tool", false))
347
+ .order("desc")), ["order", "stepOrder"]).first();
348
+ }
349
+ const addStepsArgs = {
350
+ chatId: v.id("chats"),
351
+ messageId: v.id("messages"),
352
+ steps: v.array(vStepWithMessagesWithFileAndId),
353
+ failPendingSteps: v.optional(v.boolean()),
354
+ };
355
+ export const addSteps = mutation({
356
+ args: addStepsArgs,
357
+ returns: v.array(v.doc("steps")),
358
+ handler: addStepsHandler,
359
+ });
360
+ async function addStepsHandler(ctx, args) {
361
+ const parentMessage = await ctx.db.get(args.messageId);
362
+ assert(parentMessage, `Message ${args.messageId} not found`);
363
+ const order = parentMessage.order;
364
+ assert(order !== undefined, `${args.messageId} has no order`);
365
+ let steps = await ctx.db
366
+ .query("steps")
367
+ .withIndex("parentMessageId_order_stepOrder", (q) =>
368
+ // TODO: fetch pending, and commit later
369
+ q.eq("parentMessageId", args.messageId))
370
+ .collect();
371
+ if (args.failPendingSteps) {
372
+ for (const step of steps) {
373
+ if (step.status === "pending") {
374
+ await ctx.db.patch(step._id, { status: "failed" });
375
+ }
376
+ }
377
+ steps = steps.filter((s) => s.status === "success");
378
+ }
379
+ let nextStepOrder = (steps.at(-1)?.stepOrder ?? -1) + 1;
380
+ for (const { step, messages } of args.steps) {
381
+ const stepId = await ctx.db.insert("steps", {
382
+ chatId: args.chatId,
383
+ parentMessageId: args.messageId,
384
+ order,
385
+ stepOrder: nextStepOrder,
386
+ status: step.finishReason === "stop" ? "success" : "pending",
387
+ step,
388
+ });
389
+ await addMessagesHandler(ctx, {
390
+ chatId: args.chatId,
391
+ parentMessageId: args.messageId,
392
+ stepId,
393
+ messages,
394
+ model: parentMessage.model,
395
+ agentName: parentMessage.agentName,
396
+ pending: step.finishReason === "stop" ? false : true,
397
+ failPendingSteps: false,
398
+ });
399
+ if (step.finishReason === "stop") {
400
+ await commitMessageHandler(ctx, { messageId: args.messageId });
401
+ }
402
+ steps.push((await ctx.db.get(stepId)));
403
+ nextStepOrder++;
404
+ }
405
+ return steps;
406
+ }
407
+ export const rollbackMessage = mutation({
408
+ args: {
409
+ messageId: v.id("messages"),
410
+ error: v.optional(v.string()),
411
+ },
412
+ returns: v.null(),
413
+ handler: async (ctx, { messageId, error }) => {
414
+ const message = await ctx.db.get(messageId);
415
+ assert(message, `Message ${messageId} not found`);
416
+ await ctx.db.patch(messageId, {
417
+ status: "failed",
418
+ text: error ?? message.text,
419
+ });
420
+ },
421
+ });
422
+ export const commitMessage = mutation({
423
+ args: {
424
+ messageId: v.id("messages"),
425
+ },
426
+ returns: v.null(),
427
+ handler: commitMessageHandler,
428
+ });
429
+ async function commitMessageHandler(ctx, { messageId }) {
430
+ const message = await ctx.db.get(messageId);
431
+ assert(message, `Message ${messageId} not found`);
432
+ const allSteps = await ctx.db
433
+ .query("steps")
434
+ .withIndex("parentMessageId_order_stepOrder", (q) => q.eq("parentMessageId", messageId))
435
+ .collect();
436
+ for (const step of allSteps) {
437
+ if (step.status === "pending") {
438
+ await ctx.db.patch(step._id, { status: "success" });
439
+ }
440
+ }
441
+ const order = message.order;
442
+ const messages = await mergedStream([true, false].map((tool) => stream(ctx.db, schema)
443
+ .query("messages")
444
+ .withIndex("chatId_status_tool_order_stepOrder", (q) => q
445
+ .eq("chatId", message.chatId)
446
+ .eq("status", "pending")
447
+ .eq("tool", tool)
448
+ .eq("order", order))), ["order", "stepOrder"]).collect();
449
+ for (const message of messages) {
450
+ await ctx.db.patch(message._id, { status: "success" });
451
+ }
452
+ }
453
+ export const getChatMessages = query({
454
+ args: {
455
+ chatId: v.id("chats"),
456
+ isTool: v.optional(v.boolean()),
457
+ order: v.optional(v.union(v.literal("asc"), v.literal("desc"))),
458
+ limit: v.optional(v.number()),
459
+ // Note: the other arguments cannot change from when the cursor was created.
460
+ cursor: v.optional(v.string()),
461
+ statuses: v.optional(v.array(vMessageStatus)),
462
+ },
463
+ handler: async (ctx, args) => {
464
+ const statuses = args.statuses ?? ["success"];
465
+ const toolOptions = args.isTool === undefined ? [true, false] : [args.isTool];
466
+ const order = args.order ?? "desc";
467
+ const streams = toolOptions.flatMap((tool) => statuses.map((status) => stream(ctx.db, schema)
468
+ .query("messages")
469
+ .withIndex("chatId_status_tool_order_stepOrder", (q) => q.eq("chatId", args.chatId).eq("status", status).eq("tool", tool))
470
+ .order(order)));
471
+ const messages = await mergedStream(streams, [
472
+ "order",
473
+ "stepOrder",
474
+ ]).paginate({
475
+ numItems: args.limit ?? 100,
476
+ cursor: args.cursor ?? null,
477
+ });
478
+ return {
479
+ messages: messages.page,
480
+ continueCursor: messages.continueCursor,
481
+ isDone: messages.isDone,
482
+ };
483
+ },
484
+ returns: v.object({
485
+ messages: v.array(v.doc("messages")),
486
+ continueCursor: v.string(),
487
+ isDone: v.boolean(),
488
+ }),
489
+ });
490
+ export const searchMessages = action({
491
+ args: {
492
+ userId: v.optional(v.string()),
493
+ chatId: v.optional(v.id("chats")),
494
+ ...vSearchOptions.fields,
495
+ },
496
+ returns: v.array(v.doc("messages")),
497
+ handler: async (ctx, args) => {
498
+ assert(args.userId || args.chatId, "Specify userId or chatId");
499
+ const limit = args.limit;
500
+ let textSearchMessages;
501
+ if (args.text) {
502
+ textSearchMessages = await ctx.runQuery(api.messages.textSearch, {
503
+ userId: args.userId,
504
+ chatId: args.chatId,
505
+ text: args.text,
506
+ limit,
507
+ });
508
+ }
509
+ if (args.vector) {
510
+ const dimension = args.vector.length;
511
+ if (!VectorDimensions.includes(dimension)) {
512
+ throw new Error(`Unsupported vector dimension: ${dimension}`);
513
+ }
514
+ const model = args.vectorModel ?? "unknown";
515
+ const tableName = getVectorTableName(dimension);
516
+ const vectors = (await ctx.vectorSearch(tableName, "vector", {
517
+ vector: args.vector,
518
+ filter: (q) => args.userId
519
+ ? q.eq("model_kind_userId", [model, "chat", args.userId])
520
+ : q.eq("model_kind_chatId", [model, "chat", args.chatId]),
521
+ limit,
522
+ })).filter((v) => v._score > 0.5);
523
+ // Reciprocal rank fusion
524
+ const k = 10;
525
+ const textEmbeddingIds = textSearchMessages?.map((m) => m.embeddingId);
526
+ const vectorScores = vectors
527
+ .map((v, i) => ({
528
+ id: v._id,
529
+ score: 1 / (i + k) +
530
+ 1 / (textEmbeddingIds?.indexOf(v._id) ?? Infinity + k),
531
+ }))
532
+ .sort((a, b) => b.score - a.score);
533
+ const vectorIds = vectorScores.slice(0, limit).map((v) => v.id);
534
+ const messages = await ctx.runQuery(internal.messages._fetchVectorMessages, {
535
+ userId: args.userId,
536
+ chatId: args.chatId,
537
+ vectorIds,
538
+ textSearchMessages: textSearchMessages
539
+ ?.filter((m) => !vectorIds.includes(m.embeddingId))
540
+ .slice(0, limit - vectorIds.length),
541
+ messageRange: args.messageRange ?? DEFAULT_MESSAGE_RANGE,
542
+ });
543
+ return messages;
544
+ }
545
+ return textSearchMessages?.flat() ?? [];
546
+ },
547
+ });
548
+ export const _fetchVectorMessages = internalQuery({
549
+ args: {
550
+ userId: v.optional(v.string()),
551
+ chatId: v.optional(v.id("chats")),
552
+ vectorIds: v.array(vVectorId),
553
+ textSearchMessages: v.optional(v.array(v.doc("messages"))),
554
+ messageRange: v.object({ before: v.number(), after: v.number() }),
555
+ },
556
+ returns: v.array(v.doc("messages")),
557
+ handler: async (ctx, args) => {
558
+ const messages = (await Promise.all(args.vectorIds.map((embeddingId) => ctx.db
559
+ .query("messages")
560
+ .withIndex("embeddingId", (q) => q.eq("embeddingId", embeddingId))
561
+ .filter((q) => args.userId
562
+ ? q.eq("userId", args.userId)
563
+ : // eslint-disable-next-line @typescript-eslint/no-explicit-any
564
+ q.eq("chatId", args.chatId) // not sure why it's failing...
565
+ )
566
+ .first()))).filter((m) => m !== undefined);
567
+ messages.push(...(args.textSearchMessages ?? []));
568
+ messages.sort((a, b) => a.order - b.order);
569
+ // Fetch the surrounding messages
570
+ const included = {};
571
+ for (const m of messages) {
572
+ if (!included[m.chatId]) {
573
+ included[m.chatId] = new Set();
574
+ }
575
+ included[m.chatId].add(m.order);
576
+ }
577
+ const ranges = {};
578
+ const { before, after } = args.messageRange;
579
+ for (const m of messages) {
580
+ const order = m.order;
581
+ let earliest = order - before;
582
+ let latest = order + after;
583
+ for (; earliest <= latest; earliest++) {
584
+ if (!included[m.chatId].has(earliest)) {
585
+ break;
586
+ }
587
+ }
588
+ for (; latest >= earliest; latest--) {
589
+ if (!included[m.chatId].has(latest)) {
590
+ break;
591
+ }
592
+ }
593
+ for (let i = earliest; i <= latest; i++) {
594
+ included[m.chatId].add(i);
595
+ }
596
+ if (earliest !== latest) {
597
+ const surrounding = await ctx.db
598
+ .query("messages")
599
+ .withIndex("chatId_status_tool_order_stepOrder", (q) => q
600
+ .eq("chatId", m.chatId)
601
+ .eq("status", "success")
602
+ .eq("tool", false)
603
+ .gt("order", earliest)
604
+ .lt("order", latest))
605
+ .collect();
606
+ if (!ranges[m.chatId]) {
607
+ ranges[m.chatId] = [];
608
+ }
609
+ ranges[m.chatId].push(...surrounding);
610
+ }
611
+ }
612
+ return Object.values(ranges)
613
+ .map((r) => r.sort((a, b) => a.order - b.order))
614
+ .flat();
615
+ },
616
+ });
617
+ // returns ranges of messages in order of text search relevance,
618
+ // excluding duplicates in later ranges.
619
+ export const textSearch = query({
620
+ args: {
621
+ chatId: v.optional(v.id("chats")),
622
+ userId: v.optional(v.string()),
623
+ text: v.string(),
624
+ limit: v.number(),
625
+ },
626
+ handler: async (ctx, args) => {
627
+ assert(args.userId || args.chatId, "Specify userId or chatId");
628
+ const messages = await ctx.db
629
+ .query("messages")
630
+ .withSearchIndex("text_search", (q) => args.userId
631
+ ? q.search("text", args.text).eq("userId", args.userId)
632
+ : q.search("text", args.text).eq("chatId", args.chatId))
633
+ .take(args.limit);
634
+ return messages;
635
+ },
636
+ returns: v.array(v.doc("messages")),
637
+ });
638
+ // const vMemoryConfig = v.object({
639
+ // lastMessages: v.optional(v.union(v.number(), v.literal(false))),
640
+ // semanticRecall: v.optional(
641
+ // v.union(
642
+ // v.boolean(),
643
+ // v.object({
644
+ // topK: v.number(),
645
+ // messageRange: v.union(
646
+ // v.number(),
647
+ // v.object({ before: v.number(), after: v.number() }),
648
+ // ),
649
+ // }),
650
+ // ),
651
+ // ),
652
+ // workingMemory: v.optional(
653
+ // v.object({
654
+ // enabled: v.boolean(),
655
+ // template: v.optional(v.string()),
656
+ // use: v.optional(
657
+ // v.union(v.literal("text-stream"), v.literal("tool-call")),
658
+ // ),
659
+ // }),
660
+ // ),
661
+ // threads: v.optional(
662
+ // v.object({
663
+ // generateTitle: v.optional(v.boolean()),
664
+ // }),
665
+ // ),
666
+ // });
667
+ // const vSelectBy = v.object({
668
+ // vectorSearchString: v.optional(v.string()),
669
+ // last: v.optional(v.union(v.number(), v.literal(false))),
670
+ // include: v.optional(
671
+ // v.array(
672
+ // v.object({
673
+ // id: v.string(),
674
+ // withPreviousMessages: v.optional(v.number()),
675
+ // withNextMessages: v.optional(v.number()),
676
+ // })
677
+ // )
678
+ // ),
679
+ // });
680
+ // const DEFAULT_MESSAGES_LIMIT = 40; // What pg & upstash do too.
681
+ // export const getChatMessagesPage = query({
682
+ // args: {
683
+ // threadId: v.string(),
684
+ // selectBy: v.optional(vSelectBy),
685
+ // // Unimplemented and as far I can tell no storage provider has either.
686
+ // // memoryConfig: v.optional(vMemoryConfig),
687
+ // },
688
+ // handler: async (ctx, args): Promise<SerializedMessage[]> => {
689
+ // const messages = await ctx.db
690
+ // .query("messages")
691
+ // .withIndex("threadId", (q) => q.eq("threadId", args.threadId))
692
+ // .order("desc")
693
+ // .take(args.selectBy?.last ? args.selectBy.last : DEFAULT_MESSAGES_LIMIT);
694
+ // const handled: boolean[] = [];
695
+ // const toFetch: number[] = [];
696
+ // for (const m of messages) {
697
+ // handled[m.threadOrder] = true;
698
+ // }
699
+ // await Promise.all(
700
+ // args.selectBy?.include?.map(async (range) => {
701
+ // const includeDoc = await ctx.db
702
+ // .query("messages")
703
+ // .withIndex("id", (q) => q.eq("id", range.id))
704
+ // .unique();
705
+ // if (!includeDoc) {
706
+ // console.warn(`Message ${range.id} not found`);
707
+ // return;
708
+ // }
709
+ // if (!range.withPreviousMessages && !range.withNextMessages) {
710
+ // messages.push(includeDoc);
711
+ // return;
712
+ // }
713
+ // const order = includeDoc.threadOrder;
714
+ // for (
715
+ // let i = order - (range.withPreviousMessages ?? 0);
716
+ // i < order + (range.withNextMessages ?? 0);
717
+ // i++
718
+ // ) {
719
+ // if (!handled[i]) {
720
+ // toFetch.push(i);
721
+ // handled[i] = true;
722
+ // }
723
+ // }
724
+ // }) ?? []
725
+ // );
726
+ // // sort and find unique numbers in toFetch
727
+ // const uniqueToFetch = [...new Set(toFetch)].sort();
728
+ // // find contiguous ranges in uniqueToFetch
729
+ // const ranges: { start: number; end: number }[] = [];
730
+ // for (let i = 0; i < uniqueToFetch.length; i++) {
731
+ // const start = uniqueToFetch[i];
732
+ // let end = start;
733
+ // while (i + 1 < uniqueToFetch.length && uniqueToFetch[i + 1] === end + 1) {
734
+ // end++;
735
+ // i++;
736
+ // }
737
+ // ranges.push({ start, end });
738
+ // }
739
+ // const fetched = (
740
+ // await Promise.all(
741
+ // ranges.map(async (range) => {
742
+ // return await ctx.db
743
+ // .query("messages")
744
+ // .withIndex("threadId", (q) =>
745
+ // q
746
+ // .eq("threadId", args.threadId)
747
+ // .gte("threadOrder", range.start)
748
+ // .lte("threadOrder", range.end)
749
+ // )
750
+ // .collect();
751
+ // })
752
+ // )
753
+ // ).flat();
754
+ // messages.push(...fetched);
755
+ // return messages.map(messageToSerializedMastra);
756
+ // },
757
+ // returns: v.array(vSerializedMessage),
758
+ // });
759
+ // export const saveMessages = mutation({
760
+ // args: { messages: v.array(vSerializedMessage) },
761
+ // handler: async (ctx, args) => {
762
+ // const messagesByThreadId: Record<string, SerializedMessage[]> = {};
763
+ // for (const message of args.messages) {
764
+ // messagesByThreadId[message.threadId] = [
765
+ // ...(messagesByThreadId[message.threadId] ?? []),
766
+ // message,
767
+ // ];
768
+ // }
769
+ // for (const threadId in messagesByThreadId) {
770
+ // const lastMessage = await ctx.db
771
+ // .query("messages")
772
+ // .withIndex("threadId", (q) => q.eq("threadId", threadId))
773
+ // .order("desc")
774
+ // .first();
775
+ // let threadOrder = lastMessage?.threadOrder ?? 0;
776
+ // for (const message of messagesByThreadId[threadId]) {
777
+ // threadOrder++;
778
+ // await ctx.db.insert("messages", {
779
+ // ...message,
780
+ // threadOrder,
781
+ // });
782
+ // }
783
+ // }
784
+ // },
785
+ // returns: v.null(),
786
+ // });
787
+ //# sourceMappingURL=messages.js.map