@convex-dev/agent 0.6.0 → 0.6.1

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 (184) hide show
  1. package/dist/UIMessages.d.ts +46 -0
  2. package/dist/UIMessages.d.ts.map +1 -0
  3. package/dist/UIMessages.js +634 -0
  4. package/dist/UIMessages.js.map +1 -0
  5. package/dist/client/createTool.d.ts +167 -0
  6. package/dist/client/createTool.d.ts.map +1 -0
  7. package/dist/client/createTool.js +116 -0
  8. package/dist/client/createTool.js.map +1 -0
  9. package/dist/client/defaultComponent.d.ts +11 -0
  10. package/dist/client/defaultComponent.d.ts.map +1 -0
  11. package/dist/client/defaultComponent.js +7 -0
  12. package/dist/client/defaultComponent.js.map +1 -0
  13. package/dist/client/definePlaygroundAPI.d.ts +1725 -0
  14. package/dist/client/definePlaygroundAPI.d.ts.map +1 -0
  15. package/dist/client/definePlaygroundAPI.js +271 -0
  16. package/dist/client/definePlaygroundAPI.js.map +1 -0
  17. package/dist/client/files.d.ts +69 -0
  18. package/dist/client/files.d.ts.map +1 -0
  19. package/dist/client/files.js +181 -0
  20. package/dist/client/files.js.map +1 -0
  21. package/dist/client/index.d.ts +2091 -0
  22. package/dist/client/index.d.ts.map +1 -0
  23. package/dist/client/index.js +929 -0
  24. package/dist/client/index.js.map +1 -0
  25. package/dist/client/messages.d.ts +461 -0
  26. package/dist/client/messages.d.ts.map +1 -0
  27. package/dist/client/messages.js +106 -0
  28. package/dist/client/messages.js.map +1 -0
  29. package/dist/client/mockModel.d.ts +42 -0
  30. package/dist/client/mockModel.d.ts.map +1 -0
  31. package/dist/client/mockModel.js +182 -0
  32. package/dist/client/mockModel.js.map +1 -0
  33. package/dist/client/saveInputMessages.d.ts +20 -0
  34. package/dist/client/saveInputMessages.d.ts.map +1 -0
  35. package/dist/client/saveInputMessages.js +58 -0
  36. package/dist/client/saveInputMessages.js.map +1 -0
  37. package/dist/client/search.d.ts +493 -0
  38. package/dist/client/search.d.ts.map +1 -0
  39. package/dist/client/search.js +435 -0
  40. package/dist/client/search.js.map +1 -0
  41. package/dist/client/start.d.ts +84 -0
  42. package/dist/client/start.d.ts.map +1 -0
  43. package/dist/client/start.js +185 -0
  44. package/dist/client/start.js.map +1 -0
  45. package/dist/client/streamText.d.ts +46 -0
  46. package/dist/client/streamText.d.ts.map +1 -0
  47. package/dist/client/streamText.js +127 -0
  48. package/dist/client/streamText.js.map +1 -0
  49. package/dist/client/streaming.d.ts +3778 -0
  50. package/dist/client/streaming.d.ts.map +1 -0
  51. package/dist/client/streaming.js +330 -0
  52. package/dist/client/streaming.js.map +1 -0
  53. package/dist/client/threads.d.ts +46 -0
  54. package/dist/client/threads.d.ts.map +1 -0
  55. package/dist/client/threads.js +49 -0
  56. package/dist/client/threads.js.map +1 -0
  57. package/dist/client/types.d.ts +461 -0
  58. package/dist/client/types.d.ts.map +1 -0
  59. package/dist/client/types.js +2 -0
  60. package/dist/client/types.js.map +1 -0
  61. package/dist/client/utils.d.ts +4 -0
  62. package/dist/client/utils.d.ts.map +1 -0
  63. package/dist/client/utils.js +21 -0
  64. package/dist/client/utils.js.map +1 -0
  65. package/dist/component/_generated/api.d.ts +48 -0
  66. package/dist/component/_generated/api.d.ts.map +1 -0
  67. package/dist/component/_generated/api.js +31 -0
  68. package/dist/component/_generated/api.js.map +1 -0
  69. package/dist/component/_generated/component.d.ts +3120 -0
  70. package/dist/component/_generated/component.d.ts.map +1 -0
  71. package/dist/component/_generated/component.js +11 -0
  72. package/dist/component/_generated/component.js.map +1 -0
  73. package/dist/component/_generated/dataModel.d.ts +46 -0
  74. package/dist/component/_generated/dataModel.d.ts.map +1 -0
  75. package/dist/component/_generated/dataModel.js +11 -0
  76. package/dist/component/_generated/dataModel.js.map +1 -0
  77. package/dist/component/_generated/server.d.ts +121 -0
  78. package/dist/component/_generated/server.d.ts.map +1 -0
  79. package/dist/component/_generated/server.js +78 -0
  80. package/dist/component/_generated/server.js.map +1 -0
  81. package/dist/component/apiKeys.d.ts +11 -0
  82. package/dist/component/apiKeys.d.ts.map +1 -0
  83. package/dist/component/apiKeys.js +69 -0
  84. package/dist/component/apiKeys.js.map +1 -0
  85. package/dist/component/convex.config.d.ts +3 -0
  86. package/dist/component/convex.config.d.ts.map +1 -0
  87. package/dist/component/convex.config.js +3 -0
  88. package/dist/component/convex.config.js.map +1 -0
  89. package/dist/component/files.d.ts +97 -0
  90. package/dist/component/files.d.ts.map +1 -0
  91. package/dist/component/files.js +190 -0
  92. package/dist/component/files.js.map +1 -0
  93. package/dist/component/messages.d.ts +3851 -0
  94. package/dist/component/messages.d.ts.map +1 -0
  95. package/dist/component/messages.js +757 -0
  96. package/dist/component/messages.js.map +1 -0
  97. package/dist/component/schema.d.ts +8029 -0
  98. package/dist/component/schema.d.ts.map +1 -0
  99. package/dist/component/schema.js +147 -0
  100. package/dist/component/schema.js.map +1 -0
  101. package/dist/component/streams.d.ts +128 -0
  102. package/dist/component/streams.d.ts.map +1 -0
  103. package/dist/component/streams.js +413 -0
  104. package/dist/component/streams.js.map +1 -0
  105. package/dist/component/threads.d.ts +115 -0
  106. package/dist/component/threads.d.ts.map +1 -0
  107. package/dist/component/threads.js +208 -0
  108. package/dist/component/threads.js.map +1 -0
  109. package/dist/component/users.d.ts +52 -0
  110. package/dist/component/users.d.ts.map +1 -0
  111. package/dist/component/users.js +229 -0
  112. package/dist/component/users.js.map +1 -0
  113. package/dist/component/vector/index.d.ts +61 -0
  114. package/dist/component/vector/index.d.ts.map +1 -0
  115. package/dist/component/vector/index.js +146 -0
  116. package/dist/component/vector/index.js.map +1 -0
  117. package/dist/component/vector/tables.d.ts +58 -0
  118. package/dist/component/vector/tables.d.ts.map +1 -0
  119. package/dist/component/vector/tables.js +56 -0
  120. package/dist/component/vector/tables.js.map +1 -0
  121. package/dist/deltas.d.ts +43 -0
  122. package/dist/deltas.d.ts.map +1 -0
  123. package/dist/deltas.js +446 -0
  124. package/dist/deltas.js.map +1 -0
  125. package/dist/mapping.d.ts +76 -0
  126. package/dist/mapping.d.ts.map +1 -0
  127. package/dist/mapping.js +720 -0
  128. package/dist/mapping.js.map +1 -0
  129. package/dist/react/SmoothText.d.ts +5 -0
  130. package/dist/react/SmoothText.d.ts.map +1 -0
  131. package/dist/react/SmoothText.js +6 -0
  132. package/dist/react/SmoothText.js.map +1 -0
  133. package/dist/react/index.d.ts +25 -0
  134. package/dist/react/index.d.ts.map +1 -0
  135. package/dist/react/index.js +70 -0
  136. package/dist/react/index.js.map +1 -0
  137. package/dist/react/optimisticallySendMessage.d.ts +42 -0
  138. package/dist/react/optimisticallySendMessage.d.ts.map +1 -0
  139. package/dist/react/optimisticallySendMessage.js +74 -0
  140. package/dist/react/optimisticallySendMessage.js.map +1 -0
  141. package/dist/react/types.d.ts +12 -0
  142. package/dist/react/types.d.ts.map +1 -0
  143. package/dist/react/types.js +2 -0
  144. package/dist/react/types.js.map +1 -0
  145. package/dist/react/useDeltaStreams.d.ts +10 -0
  146. package/dist/react/useDeltaStreams.d.ts.map +1 -0
  147. package/dist/react/useDeltaStreams.js +106 -0
  148. package/dist/react/useDeltaStreams.js.map +1 -0
  149. package/dist/react/useSmoothText.d.ts +27 -0
  150. package/dist/react/useSmoothText.d.ts.map +1 -0
  151. package/dist/react/useSmoothText.js +68 -0
  152. package/dist/react/useSmoothText.js.map +1 -0
  153. package/dist/react/useStreamingUIMessages.d.ts +22 -0
  154. package/dist/react/useStreamingUIMessages.d.ts.map +1 -0
  155. package/dist/react/useStreamingUIMessages.js +92 -0
  156. package/dist/react/useStreamingUIMessages.js.map +1 -0
  157. package/dist/react/useThreadMessages.d.ts +104 -0
  158. package/dist/react/useThreadMessages.d.ts.map +1 -0
  159. package/dist/react/useThreadMessages.js +148 -0
  160. package/dist/react/useThreadMessages.js.map +1 -0
  161. package/dist/react/useUIMessages.d.ts +96 -0
  162. package/dist/react/useUIMessages.d.ts.map +1 -0
  163. package/dist/react/useUIMessages.js +108 -0
  164. package/dist/react/useUIMessages.js.map +1 -0
  165. package/dist/shared.d.ts +26 -0
  166. package/dist/shared.d.ts.map +1 -0
  167. package/dist/shared.js +67 -0
  168. package/dist/shared.js.map +1 -0
  169. package/dist/validators.d.ts +24516 -0
  170. package/dist/validators.d.ts.map +1 -0
  171. package/dist/validators.js +475 -0
  172. package/dist/validators.js.map +1 -0
  173. package/package.json +39 -39
  174. package/src/UIMessages.ts +0 -9
  175. package/src/client/approval.test.ts +144 -0
  176. package/src/client/index.ts +17 -9
  177. package/src/client/search.ts +0 -29
  178. package/src/client/start.ts +2 -2
  179. package/src/client/streamText.ts +9 -0
  180. package/src/client/streaming.ts +30 -23
  181. package/src/component/_generated/dataModel.ts +1 -1
  182. package/src/component/_generated/server.ts +0 -5
  183. package/src/mapping.ts +38 -22
  184. package/src/node_modules/.vite/vitest/da39a3ee5e6b4b0d3255bfef95601890afd80709/results.json +1 -0
@@ -0,0 +1,757 @@
1
+ import { assert, omit, pick } from "convex-helpers";
2
+ import { mergedStream, stream } from "convex-helpers/server/stream";
3
+ import { paginationOptsValidator, } from "convex/server";
4
+ import { DEFAULT_MESSAGE_RANGE, DEFAULT_RECENT_MESSAGES, extractText, isTool, sorted, } from "../shared.js";
5
+ import { vMessageDoc, vMessageEmbeddingsWithDimension, vMessageStatus, vMessageWithMetadataInternal, vPaginationResult, } from "../validators.js";
6
+ import { api, internal } from "./_generated/api.js";
7
+ import { action, internalMutation, internalQuery, mutation, query, } from "./_generated/server.js";
8
+ import { schema, v } from "./schema.js";
9
+ import { insertVector, searchVectors } from "./vector/index.js";
10
+ import { validateVectorDimension, vVectorId, } from "./vector/tables.js";
11
+ import { changeRefcount } from "./files.js";
12
+ import { getStreamingMessagesWithMetadata, finishHandler } from "./streams.js";
13
+ import { partial } from "convex-helpers/validators";
14
+ function publicMessage(message) {
15
+ return omit(message, ["parentMessageId", "stepId", "files"]);
16
+ }
17
+ export async function deleteMessage(ctx, messageDoc) {
18
+ await ctx.db.delete(messageDoc._id);
19
+ if (messageDoc.embeddingId) {
20
+ await ctx.db.delete(messageDoc.embeddingId);
21
+ }
22
+ if (messageDoc.fileIds) {
23
+ await changeRefcount(ctx, messageDoc.fileIds, []);
24
+ }
25
+ }
26
+ export const deleteByIds = mutation({
27
+ args: { messageIds: v.array(v.id("messages")) },
28
+ returns: v.array(v.id("messages")),
29
+ handler: async (ctx, args) => {
30
+ const deletedMessageIds = await Promise.all(args.messageIds.map(async (id) => {
31
+ const message = await ctx.db.get(id);
32
+ if (message) {
33
+ await deleteMessage(ctx, message);
34
+ return id;
35
+ }
36
+ return null;
37
+ }));
38
+ return deletedMessageIds.filter((id) => id !== null);
39
+ },
40
+ });
41
+ export const messageStatuses = vMessageDoc.fields.status.members.map((m) => m.value);
42
+ export const deleteByOrder = mutation({
43
+ args: {
44
+ threadId: v.id("threads"),
45
+ startOrder: v.number(),
46
+ startStepOrder: v.optional(v.number()),
47
+ endOrder: v.number(),
48
+ endStepOrder: v.optional(v.number()),
49
+ },
50
+ returns: v.object({
51
+ isDone: v.boolean(),
52
+ lastOrder: v.optional(v.number()),
53
+ lastStepOrder: v.optional(v.number()),
54
+ }),
55
+ handler: async (ctx, args) => {
56
+ const messages = await orderedMessagesStream(ctx, {
57
+ threadId: args.threadId,
58
+ sortOrder: "asc",
59
+ startOrder: args.startOrder,
60
+ startOrderBound: "gte",
61
+ })
62
+ .narrow({
63
+ lowerBound: args.startStepOrder
64
+ ? [args.startOrder, args.startStepOrder]
65
+ : [args.startOrder],
66
+ lowerBoundInclusive: true,
67
+ upperBound: args.endStepOrder
68
+ ? [args.endOrder, args.endStepOrder]
69
+ : [args.endOrder],
70
+ upperBoundInclusive: false,
71
+ })
72
+ .take(64);
73
+ await Promise.all(messages.map((m) => deleteMessage(ctx, m)));
74
+ return {
75
+ isDone: messages.length < 64,
76
+ lastOrder: messages.at(-1)?.order,
77
+ lastStepOrder: messages.at(-1)?.stepOrder,
78
+ };
79
+ },
80
+ });
81
+ const addMessagesArgs = {
82
+ userId: v.optional(v.string()),
83
+ threadId: v.id("threads"),
84
+ promptMessageId: v.optional(v.id("messages")),
85
+ agentName: v.optional(v.string()),
86
+ messages: v.array(vMessageWithMetadataInternal),
87
+ embeddings: v.optional(vMessageEmbeddingsWithDimension),
88
+ failPendingSteps: v.optional(v.boolean()),
89
+ // A pending message to update. If the pending message failed, abort.
90
+ pendingMessageId: v.optional(v.id("messages")),
91
+ // if set to true, these messages will not show up in text or vector search
92
+ // results for the userId
93
+ hideFromUserIdSearch: v.optional(v.boolean()),
94
+ // If provided, finish this stream atomically with the message save.
95
+ // This prevents UI flickering from separate mutations (issue #181).
96
+ finishStreamId: v.optional(v.id("streamingMessages")),
97
+ };
98
+ export const addMessages = mutation({
99
+ args: addMessagesArgs,
100
+ handler: addMessagesHandler,
101
+ returns: v.object({ messages: v.array(vMessageDoc) }),
102
+ });
103
+ async function addMessagesHandler(ctx, args) {
104
+ let userId = args.userId;
105
+ const threadId = args.threadId;
106
+ if (!userId && args.threadId) {
107
+ const thread = await ctx.db.get(args.threadId);
108
+ assert(thread, `Thread ${args.threadId} not found`);
109
+ userId = thread.userId;
110
+ }
111
+ const { embeddings, failPendingSteps,
112
+ // Destructured separately to exclude from `...rest` (used in addMessages args, not message fields)
113
+ finishStreamId, messages, promptMessageId, pendingMessageId, hideFromUserIdSearch, ...rest } = args;
114
+ const promptMessage = promptMessageId && (await ctx.db.get(promptMessageId));
115
+ if (failPendingSteps) {
116
+ assert(args.threadId, "threadId is required to fail pending steps");
117
+ const pendingMessages = await ctx.db
118
+ .query("messages")
119
+ .withIndex("threadId_status_tool_order_stepOrder", (q) => q.eq("threadId", threadId).eq("status", "pending"))
120
+ .order("desc")
121
+ .take(100);
122
+ await Promise.all(pendingMessages
123
+ .filter((m) => !promptMessage || m.order === promptMessage.order)
124
+ .filter((m) => !pendingMessageId || m._id !== pendingMessageId)
125
+ .map(async (m) => {
126
+ if (m.embeddingId) {
127
+ await ctx.db.delete(m.embeddingId);
128
+ }
129
+ await ctx.db.patch(m._id, {
130
+ status: "failed",
131
+ error: "Restarting",
132
+ embeddingId: undefined,
133
+ });
134
+ }));
135
+ }
136
+ let order, stepOrder;
137
+ let fail = false;
138
+ let error;
139
+ if (promptMessageId) {
140
+ assert(promptMessage, `Parent message ${promptMessageId} not found`);
141
+ if (promptMessage.status === "failed") {
142
+ fail = true;
143
+ error = promptMessage.error ?? error ?? "The prompt message failed";
144
+ }
145
+ order = promptMessage.order;
146
+ // Defend against there being existing messages with this parent.
147
+ const maxMessage = await getMaxMessage(ctx, threadId, order);
148
+ stepOrder = maxMessage?.stepOrder ?? promptMessage.stepOrder;
149
+ }
150
+ else {
151
+ const maxMessage = await getMaxMessage(ctx, threadId);
152
+ order = maxMessage?.order ?? -1;
153
+ stepOrder = maxMessage?.stepOrder ?? -1;
154
+ }
155
+ const toReturn = [];
156
+ if (embeddings) {
157
+ assert(embeddings.vectors.length === messages.length, "embeddings.vectors.length must match messages.length");
158
+ }
159
+ for (let i = 0; i < messages.length; i++) {
160
+ const message = messages[i];
161
+ let embeddingId;
162
+ if (embeddings &&
163
+ embeddings.vectors[i] &&
164
+ !fail &&
165
+ message.status !== "failed") {
166
+ embeddingId = await insertVector(ctx, embeddings.dimension, {
167
+ vector: embeddings.vectors[i],
168
+ model: embeddings.model,
169
+ table: "messages",
170
+ userId: hideFromUserIdSearch ? undefined : userId,
171
+ threadId,
172
+ });
173
+ }
174
+ const messageDoc = {
175
+ ...rest,
176
+ ...message,
177
+ embeddingId,
178
+ parentMessageId: promptMessageId,
179
+ userId,
180
+ tool: isTool(message.message),
181
+ text: hideFromUserIdSearch ? undefined : extractText(message.message),
182
+ status: fail ? "failed" : (message.status ?? "success"),
183
+ error: fail ? error : message.error,
184
+ };
185
+ // If there is a pending message, we replace that one with the first message
186
+ // and subsequent ones will follow the regular order/subOrder advancement.
187
+ if (i === 0 && pendingMessageId) {
188
+ const pendingMessage = await ctx.db.get(pendingMessageId);
189
+ assert(pendingMessage, `Pending msg ${pendingMessageId} not found`);
190
+ if (pendingMessage.status === "failed") {
191
+ fail = true;
192
+ error =
193
+ `Trying to update a message that failed: ${pendingMessageId}, ` +
194
+ `error: ${pendingMessage.error ?? error}`;
195
+ messageDoc.status = "failed";
196
+ messageDoc.error = error;
197
+ }
198
+ if (message.fileIds) {
199
+ await changeRefcount(ctx, pendingMessage.fileIds ?? [], message.fileIds);
200
+ }
201
+ await ctx.db.replace(pendingMessage._id, {
202
+ ...messageDoc,
203
+ order: pendingMessage.order,
204
+ stepOrder: pendingMessage.stepOrder,
205
+ });
206
+ toReturn.push((await ctx.db.get(pendingMessage._id)));
207
+ continue;
208
+ }
209
+ if (message.message.role === "user") {
210
+ if (promptMessage && promptMessage.order === order) {
211
+ // see if there's a later message than the parent message order
212
+ const maxMessage = await getMaxMessage(ctx, threadId);
213
+ order = (maxMessage?.order ?? order) + 1;
214
+ }
215
+ else {
216
+ order++;
217
+ }
218
+ stepOrder = 0;
219
+ }
220
+ else {
221
+ if (order < 0) {
222
+ order = 0;
223
+ }
224
+ stepOrder++;
225
+ }
226
+ const messageId = await ctx.db.insert("messages", {
227
+ ...messageDoc,
228
+ order,
229
+ stepOrder,
230
+ });
231
+ if (message.fileIds) {
232
+ await changeRefcount(ctx, [], message.fileIds);
233
+ }
234
+ // TODO: delete the associated stream data for the order/stepOrder
235
+ toReturn.push((await ctx.db.get(messageId)));
236
+ }
237
+ // Atomically finish the stream if requested, preventing UI flickering
238
+ // from separate mutations for message save and stream finish (issue #181).
239
+ if (finishStreamId) {
240
+ await finishHandler(ctx, { streamId: finishStreamId });
241
+ }
242
+ return { messages: toReturn.map(publicMessage) };
243
+ }
244
+ // exported for tests
245
+ export async function getMaxMessage(ctx, threadId, order) {
246
+ return orderedMessagesStream(ctx, {
247
+ threadId,
248
+ sortOrder: "desc",
249
+ startOrder: order,
250
+ startOrderBound: "eq",
251
+ }).first();
252
+ }
253
+ function orderedMessagesStream(ctx, args) {
254
+ return mergedStream([true, false].flatMap((tool) => messageStatuses.map((status) => stream(ctx.db, schema)
255
+ .query("messages")
256
+ .withIndex("threadId_status_tool_order_stepOrder", (q) => {
257
+ const qq = q
258
+ .eq("threadId", args.threadId)
259
+ .eq("status", status)
260
+ .eq("tool", tool);
261
+ if (args.startOrder !== undefined) {
262
+ if (args.startOrderBound === "gte") {
263
+ return qq.gte("order", args.startOrder);
264
+ }
265
+ else {
266
+ return qq.eq("order", args.startOrder);
267
+ }
268
+ }
269
+ return qq;
270
+ })
271
+ .order(args.sortOrder))), ["order", "stepOrder"]);
272
+ }
273
+ export const finalizeMessage = mutation({
274
+ args: {
275
+ messageId: v.id("messages"),
276
+ result: v.union(v.object({ status: v.literal("success") }), v.object({ status: v.literal("failed"), error: v.string() })),
277
+ },
278
+ returns: v.null(),
279
+ handler: async (ctx, { messageId, result }) => {
280
+ const message = await ctx.db.get(messageId);
281
+ assert(message, `Message ${messageId} not found`);
282
+ if (message.status !== "pending") {
283
+ console.debug("Trying to finalize a message that's already", message.status);
284
+ return;
285
+ }
286
+ // See if we can add any in-progress data
287
+ if (!message.message?.content.length) {
288
+ const messages = await getStreamingMessagesWithMetadata(ctx, message, result);
289
+ if (messages.length > 0) {
290
+ await addMessagesHandler(ctx, {
291
+ messages,
292
+ threadId: message.threadId,
293
+ agentName: message.agentName,
294
+ failPendingSteps: false,
295
+ pendingMessageId: messageId,
296
+ userId: message.userId,
297
+ embeddings: undefined,
298
+ });
299
+ return;
300
+ }
301
+ }
302
+ if (result.status === "failed") {
303
+ if (message.embeddingId) {
304
+ await ctx.db.delete(message.embeddingId);
305
+ }
306
+ await ctx.db.patch(messageId, {
307
+ status: "failed",
308
+ error: result.error,
309
+ embeddingId: undefined,
310
+ });
311
+ }
312
+ else {
313
+ await ctx.db.patch(messageId, { status: "success" });
314
+ }
315
+ },
316
+ });
317
+ export const updateMessage = mutation({
318
+ args: {
319
+ messageId: v.id("messages"),
320
+ patch: v.object(partial(pick(schema.tables.messages.validator.fields, [
321
+ "message",
322
+ "fileIds",
323
+ "status",
324
+ "error",
325
+ "model",
326
+ "provider",
327
+ "providerOptions",
328
+ "finishReason",
329
+ ]))),
330
+ },
331
+ returns: vMessageDoc,
332
+ handler: async (ctx, args) => {
333
+ const message = await ctx.db.get(args.messageId);
334
+ assert(message, `Message ${args.messageId} not found`);
335
+ if (args.patch.fileIds) {
336
+ await changeRefcount(ctx, message.fileIds ?? [], args.patch.fileIds);
337
+ }
338
+ const patch = { ...args.patch };
339
+ if (args.patch.message !== undefined) {
340
+ patch.message = args.patch.message;
341
+ patch.tool = isTool(args.patch.message);
342
+ patch.text = extractText(args.patch.message);
343
+ }
344
+ if (args.patch.status === "failed") {
345
+ if (message.embeddingId) {
346
+ await ctx.db.delete(message.embeddingId);
347
+ }
348
+ patch.embeddingId = undefined;
349
+ }
350
+ await ctx.db.patch(args.messageId, patch);
351
+ return publicMessage((await ctx.db.get(args.messageId)));
352
+ },
353
+ });
354
+ const cloneMessageArgs = {
355
+ sourceThreadId: v.id("threads"),
356
+ targetThreadId: v.id("threads"),
357
+ // defaults to false, so searching for a message by userId will not find
358
+ // these copies
359
+ copyUserIdForVectorSearch: v.optional(v.boolean()),
360
+ // defaults to false, so tool calls & responses will be copied
361
+ excludeToolMessages: v.optional(v.boolean()),
362
+ // defaults to copying all messages, but you could just copy success messages.
363
+ statuses: v.optional(v.array(vMessageStatus)),
364
+ // stop at this message id
365
+ upToAndIncludingMessageId: v.optional(v.id("messages")),
366
+ // defaults to 0. the messages will be inserted starting at this order.
367
+ insertAtOrder: v.optional(v.number()),
368
+ };
369
+ export const cloneMessageBatch = internalMutation({
370
+ args: {
371
+ ...cloneMessageArgs,
372
+ paginationOpts: paginationOptsValidator,
373
+ },
374
+ handler: async (ctx, args) => {
375
+ const orderOffset = args.insertAtOrder ?? 0;
376
+ const result = await listMessagesByThreadIdHandler(ctx, {
377
+ threadId: args.sourceThreadId,
378
+ excludeToolMessages: args.excludeToolMessages,
379
+ order: "desc",
380
+ paginationOpts: args.paginationOpts,
381
+ statuses: args.statuses,
382
+ upToAndIncludingMessageId: args.upToAndIncludingMessageId,
383
+ });
384
+ const existing = result.page.length === 0
385
+ ? []
386
+ : await mergedStream([true, false].flatMap((tool) => messageStatuses.map((status) => stream(ctx.db, schema)
387
+ .query("messages")
388
+ .withIndex("threadId_status_tool_order_stepOrder", (q) => q
389
+ .eq("threadId", args.targetThreadId)
390
+ .eq("status", status)
391
+ .eq("tool", tool)
392
+ .gte("order", result.page[0].order)
393
+ .lte("order", result.page[result.page.length - 1].order)))), ["order", "stepOrder"]).collect();
394
+ await Promise.all(result.page
395
+ .filter((m) => !existing.some((e) => e.order === m.order && e.stepOrder === m.stepOrder))
396
+ .map(async (m) => {
397
+ // update file refs
398
+ if (m.fileIds) {
399
+ await changeRefcount(ctx, [], m.fileIds);
400
+ }
401
+ let embeddingId = undefined;
402
+ if (m.embeddingId) {
403
+ const vector = await ctx.db.get(m.embeddingId);
404
+ assert(vector, `Vector ${m.embeddingId} not found`);
405
+ const dimension = vector.vector.length;
406
+ validateVectorDimension(dimension);
407
+ embeddingId = await insertVector(ctx, dimension, {
408
+ ...pick(vector, ["model", "table", "vector"]),
409
+ userId: args.copyUserIdForVectorSearch
410
+ ? vector.userId
411
+ : undefined,
412
+ threadId: args.targetThreadId,
413
+ });
414
+ }
415
+ await ctx.db.insert("messages", {
416
+ ...omit(m, [
417
+ "_id",
418
+ "_creationTime",
419
+ "threadId",
420
+ "order",
421
+ "embeddingId",
422
+ ]),
423
+ embeddingId,
424
+ threadId: args.targetThreadId,
425
+ order: orderOffset + m.order,
426
+ });
427
+ }));
428
+ return {
429
+ numCopied: result.page.length,
430
+ continueCursor: result.continueCursor,
431
+ isDone: result.isDone,
432
+ };
433
+ },
434
+ });
435
+ export const cloneThread = action({
436
+ args: {
437
+ ...cloneMessageArgs,
438
+ batchSize: v.optional(v.number()),
439
+ // how many messages to copy
440
+ limit: v.optional(v.number()),
441
+ },
442
+ returns: v.number(),
443
+ handler: async (ctx, args) => {
444
+ let cursor = null;
445
+ let copiedSoFar = 0;
446
+ while (copiedSoFar < (args.limit ?? Infinity)) {
447
+ const numToCopy = Math.min(args.batchSize ?? DEFAULT_RECENT_MESSAGES, args.limit ?? Infinity - copiedSoFar);
448
+ const result = await ctx.runMutation(internal.messages.cloneMessageBatch, {
449
+ ...args,
450
+ paginationOpts: {
451
+ cursor,
452
+ numItems: numToCopy,
453
+ },
454
+ });
455
+ copiedSoFar += result.numCopied;
456
+ cursor = result.continueCursor;
457
+ if (result.isDone) {
458
+ break;
459
+ }
460
+ }
461
+ return copiedSoFar;
462
+ },
463
+ });
464
+ export const listMessagesByThreadIdArgs = {
465
+ threadId: v.id("threads"),
466
+ excludeToolMessages: v.optional(v.boolean()),
467
+ /** What order to sort the messages in. To get the latest, use "desc". */
468
+ order: v.union(v.literal("asc"), v.literal("desc")),
469
+ paginationOpts: v.optional(paginationOptsValidator),
470
+ statuses: v.optional(v.array(vMessageStatus)),
471
+ upToAndIncludingMessageId: v.optional(v.id("messages")),
472
+ };
473
+ export const listMessagesByThreadId = query({
474
+ args: listMessagesByThreadIdArgs,
475
+ handler: async (ctx, args) => {
476
+ const messages = await listMessagesByThreadIdHandler(ctx, args);
477
+ return { ...messages, page: messages.page.map(publicMessage) };
478
+ },
479
+ returns: vPaginationResult(vMessageDoc),
480
+ });
481
+ async function listMessagesByThreadIdHandler(ctx, args) {
482
+ const statuses = args.statuses ?? vMessageStatus.members.map((m) => m.value);
483
+ const last = args.upToAndIncludingMessageId &&
484
+ (await ctx.db.get(args.upToAndIncludingMessageId));
485
+ assert(!last || last.threadId === args.threadId, "upToAndIncludingMessageId must be a message in the thread");
486
+ const toolOptions = args.excludeToolMessages ? [false] : [true, false];
487
+ const order = args.order ?? "desc";
488
+ const streams = toolOptions.flatMap((tool) => statuses.map((status) => stream(ctx.db, schema)
489
+ .query("messages")
490
+ .withIndex("threadId_status_tool_order_stepOrder", (q) => {
491
+ const qq = q
492
+ .eq("threadId", args.threadId)
493
+ .eq("status", status)
494
+ .eq("tool", tool);
495
+ if (last) {
496
+ return qq.lte("order", last.order);
497
+ }
498
+ return qq;
499
+ })
500
+ .order(order)
501
+ .filterWith(
502
+ // We allow all messages on the same order.
503
+ async (m) => !last || m.order <= last.order)));
504
+ const messages = await mergedStream(streams, ["order", "stepOrder"]).paginate(args.paginationOpts ?? {
505
+ numItems: DEFAULT_RECENT_MESSAGES,
506
+ cursor: null,
507
+ });
508
+ if (messages.page.length === 0) {
509
+ messages.isDone = true;
510
+ }
511
+ return messages;
512
+ }
513
+ export const getMessagesByIds = query({
514
+ args: { messageIds: v.array(v.id("messages")) },
515
+ handler: async (ctx, args) => {
516
+ return (await Promise.all(args.messageIds.map((id) => ctx.db.get(id)))).map((m) => (m ? publicMessage(m) : null));
517
+ },
518
+ returns: v.array(v.union(v.null(), vMessageDoc)),
519
+ });
520
+ export const searchMessages = action({
521
+ args: {
522
+ threadId: v.optional(v.id("threads")),
523
+ searchAllMessagesForUserId: v.optional(v.string()),
524
+ targetMessageId: v.optional(v.id("messages")),
525
+ embedding: v.optional(v.array(v.number())),
526
+ embeddingModel: v.optional(v.string()),
527
+ text: v.optional(v.string()),
528
+ textSearch: v.optional(v.boolean()),
529
+ vectorSearch: v.optional(v.boolean()),
530
+ limit: v.number(),
531
+ vectorScoreThreshold: v.optional(v.number()),
532
+ messageRange: v.optional(v.object({ before: v.number(), after: v.number() })),
533
+ },
534
+ returns: v.array(vMessageDoc),
535
+ handler: async (ctx, args) => {
536
+ assert(args.searchAllMessagesForUserId || args.threadId, "Specify userId or threadId");
537
+ const limit = args.limit;
538
+ let textSearchMessages;
539
+ if (args.textSearch) {
540
+ textSearchMessages = await ctx.runQuery(api.messages.textSearch, {
541
+ searchAllMessagesForUserId: args.searchAllMessagesForUserId,
542
+ threadId: args.threadId,
543
+ targetMessageId: args.targetMessageId,
544
+ text: args.text,
545
+ limit,
546
+ });
547
+ }
548
+ if (args.vectorSearch) {
549
+ let embedding = args.embedding;
550
+ let model = args.embeddingModel;
551
+ if (!embedding) {
552
+ if (args.targetMessageId) {
553
+ const target = await ctx.runQuery(api.messages.getMessageSearchFields, {
554
+ messageId: args.targetMessageId,
555
+ });
556
+ assert(target, "Target message embedding not found.");
557
+ embedding = target.embedding;
558
+ model = target.embeddingModel;
559
+ }
560
+ }
561
+ assert(embedding && model, "Embedding missing");
562
+ const dimension = embedding.length;
563
+ validateVectorDimension(dimension);
564
+ const vectors = (await searchVectors(ctx, embedding, {
565
+ dimension,
566
+ model,
567
+ table: "messages",
568
+ searchAllMessagesForUserId: args.searchAllMessagesForUserId,
569
+ threadId: args.threadId,
570
+ limit,
571
+ })).filter((v) => v._score > (args.vectorScoreThreshold ?? 0));
572
+ // Reciprocal rank fusion
573
+ const k = 10;
574
+ const textEmbeddingIds = textSearchMessages?.map((m) => m.embeddingId);
575
+ const vectorScores = vectors
576
+ .map((v, i) => ({
577
+ id: v._id,
578
+ score: 1 / (i + k) +
579
+ 1 / ((textEmbeddingIds?.indexOf(v._id) ?? Infinity) + k),
580
+ }))
581
+ .sort((a, b) => b.score - a.score);
582
+ const embeddingIds = vectorScores.slice(0, limit).map((v) => v.id);
583
+ const messages = await ctx.runQuery(internal.messages._fetchSearchMessages, {
584
+ searchAllMessagesForUserId: args.searchAllMessagesForUserId,
585
+ threadId: args.threadId,
586
+ embeddingIds,
587
+ textSearchMessages: textSearchMessages?.filter((m) => !embeddingIds.includes(m.embeddingId)),
588
+ messageRange: args.messageRange ?? DEFAULT_MESSAGE_RANGE,
589
+ beforeMessageId: args.targetMessageId,
590
+ limit,
591
+ });
592
+ return messages;
593
+ }
594
+ return textSearchMessages?.flat() ?? [];
595
+ },
596
+ });
597
+ export const _fetchSearchMessages = internalQuery({
598
+ args: {
599
+ threadId: v.optional(v.id("threads")),
600
+ embeddingIds: v.array(vVectorId),
601
+ searchAllMessagesForUserId: v.optional(v.string()),
602
+ textSearchMessages: v.optional(v.array(vMessageDoc)),
603
+ messageRange: v.object({ before: v.number(), after: v.number() }),
604
+ beforeMessageId: v.optional(v.id("messages")),
605
+ limit: v.number(),
606
+ },
607
+ returns: v.array(vMessageDoc),
608
+ handler: async (ctx, args) => {
609
+ const beforeMessage = args.beforeMessageId && (await ctx.db.get(args.beforeMessageId));
610
+ const { searchAllMessagesForUserId, threadId } = args;
611
+ assert(searchAllMessagesForUserId || threadId, "Specify searchAllMessagesForUserId or threadId to search");
612
+ let messages = (await Promise.all(args.embeddingIds.map((embeddingId) => ctx.db
613
+ .query("messages")
614
+ .withIndex("embeddingId_threadId", (q) => searchAllMessagesForUserId
615
+ ? q.eq("embeddingId", embeddingId)
616
+ : q.eq("embeddingId", embeddingId).eq("threadId", threadId))
617
+ .filter((q) => q.and(q.eq(q.field("status"), "success"), searchAllMessagesForUserId
618
+ ? q.eq(q.field("userId"), searchAllMessagesForUserId)
619
+ : q.eq(q.field("threadId"), threadId)))
620
+ .first())))
621
+ .filter((m) => m !== undefined &&
622
+ m !== null &&
623
+ !m.tool &&
624
+ (!beforeMessage ||
625
+ m.order < beforeMessage.order ||
626
+ (m.order === beforeMessage.order &&
627
+ m.stepOrder < beforeMessage.stepOrder)))
628
+ .map(publicMessage);
629
+ messages.push(...(args.textSearchMessages ?? []));
630
+ // TODO: prioritize more recent messages
631
+ messages = sorted(messages);
632
+ messages = messages.slice(0, args.limit);
633
+ // Fetch the surrounding messages
634
+ if (!threadId) {
635
+ return messages;
636
+ }
637
+ const included = {};
638
+ for (const m of messages) {
639
+ const searchId = m.threadId ?? m.userId;
640
+ if (!included[searchId]) {
641
+ included[searchId] = new Set();
642
+ }
643
+ included[searchId].add(m.order);
644
+ }
645
+ const ranges = {};
646
+ const { before, after } = args.messageRange;
647
+ for (const m of messages) {
648
+ const searchId = m.threadId ?? m.userId;
649
+ const order = m.order;
650
+ let earliest = order - before;
651
+ let latest = order + after;
652
+ for (; earliest <= latest; earliest++) {
653
+ if (!included[searchId].has(earliest)) {
654
+ break;
655
+ }
656
+ }
657
+ for (; latest >= earliest; latest--) {
658
+ if (!included[searchId].has(latest)) {
659
+ break;
660
+ }
661
+ }
662
+ for (let i = earliest; i <= latest; i++) {
663
+ included[searchId].add(i);
664
+ }
665
+ if (earliest !== latest) {
666
+ const surrounding = await ctx.db
667
+ .query("messages")
668
+ .withIndex("threadId_status_tool_order_stepOrder", (q) => q
669
+ .eq("threadId", m.threadId)
670
+ .eq("status", "success")
671
+ .eq("tool", false)
672
+ .gte("order", earliest)
673
+ .lte("order", latest))
674
+ .collect();
675
+ if (!ranges[searchId]) {
676
+ ranges[searchId] = [];
677
+ }
678
+ ranges[searchId].push(...surrounding);
679
+ }
680
+ }
681
+ for (const r of Object.values(ranges).flat()) {
682
+ if (!messages.some((m) => m._id === r._id)) {
683
+ messages.push(publicMessage(r));
684
+ }
685
+ }
686
+ return sorted(messages);
687
+ },
688
+ });
689
+ // returns ranges of messages in order of text search relevance,
690
+ // excluding duplicates in later ranges.
691
+ export const textSearch = query({
692
+ args: {
693
+ threadId: v.optional(v.id("threads")),
694
+ searchAllMessagesForUserId: v.optional(v.string()),
695
+ text: v.optional(v.string()),
696
+ targetMessageId: v.optional(v.id("messages")),
697
+ limit: v.number(),
698
+ },
699
+ handler: async (ctx, args) => {
700
+ assert(args.searchAllMessagesForUserId || args.threadId, "Specify userId or threadId");
701
+ const targetMessage = args.targetMessageId && (await ctx.db.get(args.targetMessageId));
702
+ const order = targetMessage?.order;
703
+ const text = args.text || targetMessage?.text;
704
+ if (!text) {
705
+ console.warn("No text to search", targetMessage, args.text);
706
+ return [];
707
+ }
708
+ const messages = await ctx.db
709
+ .query("messages")
710
+ .withSearchIndex("text_search", (q) => args.searchAllMessagesForUserId
711
+ ? q.search("text", text).eq("userId", args.searchAllMessagesForUserId)
712
+ : q.search("text", text).eq("threadId", args.threadId))
713
+ // Just in case tool messages slip through
714
+ .filter((q) => {
715
+ const qq = q.eq(q.field("tool"), false);
716
+ if (order) {
717
+ return q.and(qq, q.lte(q.field("order"), order));
718
+ }
719
+ return qq;
720
+ })
721
+ .take(args.limit);
722
+ return messages
723
+ .filter((m) => !targetMessage ||
724
+ m.order < targetMessage.order ||
725
+ (m.order === targetMessage.order &&
726
+ m.stepOrder < targetMessage.stepOrder))
727
+ .map(publicMessage);
728
+ },
729
+ returns: v.array(vMessageDoc),
730
+ });
731
+ export const getMessageSearchFields = query({
732
+ args: {
733
+ messageId: v.id("messages"),
734
+ },
735
+ returns: v.object({
736
+ text: v.optional(v.string()),
737
+ embedding: v.optional(v.array(v.number())),
738
+ embeddingModel: v.optional(v.string()),
739
+ }),
740
+ handler: async (ctx, args) => {
741
+ const message = await ctx.db.get(args.messageId);
742
+ const text = message?.text;
743
+ let embedding = undefined;
744
+ let embeddingModel = undefined;
745
+ if (message?.embeddingId) {
746
+ const target = await ctx.db.get(message.embeddingId);
747
+ embedding = target?.vector;
748
+ embeddingModel = target?.model;
749
+ }
750
+ return {
751
+ text,
752
+ embedding,
753
+ embeddingModel,
754
+ };
755
+ },
756
+ });
757
+ //# sourceMappingURL=messages.js.map