@convex-dev/agent 0.6.0-alpha.1 → 0.6.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 (181) hide show
  1. package/package.json +1 -1
  2. package/src/UIMessages.ts +135 -0
  3. package/src/client/index.ts +56 -14
  4. package/src/client/search.test.ts +4 -5
  5. package/src/client/search.ts +45 -4
  6. package/src/client/streaming.integration.test.ts +1206 -0
  7. package/src/client/streaming.ts +9 -2
  8. package/src/mapping.test.ts +136 -71
  9. package/src/mapping.ts +81 -28
  10. package/dist/UIMessages.d.ts +0 -46
  11. package/dist/UIMessages.d.ts.map +0 -1
  12. package/dist/UIMessages.js +0 -546
  13. package/dist/UIMessages.js.map +0 -1
  14. package/dist/client/createTool.d.ts +0 -167
  15. package/dist/client/createTool.d.ts.map +0 -1
  16. package/dist/client/createTool.js +0 -116
  17. package/dist/client/createTool.js.map +0 -1
  18. package/dist/client/defaultComponent.d.ts +0 -11
  19. package/dist/client/defaultComponent.d.ts.map +0 -1
  20. package/dist/client/defaultComponent.js +0 -7
  21. package/dist/client/defaultComponent.js.map +0 -1
  22. package/dist/client/definePlaygroundAPI.d.ts +0 -1725
  23. package/dist/client/definePlaygroundAPI.d.ts.map +0 -1
  24. package/dist/client/definePlaygroundAPI.js +0 -271
  25. package/dist/client/definePlaygroundAPI.js.map +0 -1
  26. package/dist/client/files.d.ts +0 -69
  27. package/dist/client/files.d.ts.map +0 -1
  28. package/dist/client/files.js +0 -181
  29. package/dist/client/files.js.map +0 -1
  30. package/dist/client/index.d.ts +0 -2091
  31. package/dist/client/index.d.ts.map +0 -1
  32. package/dist/client/index.js +0 -895
  33. package/dist/client/index.js.map +0 -1
  34. package/dist/client/messages.d.ts +0 -461
  35. package/dist/client/messages.d.ts.map +0 -1
  36. package/dist/client/messages.js +0 -106
  37. package/dist/client/messages.js.map +0 -1
  38. package/dist/client/mockModel.d.ts +0 -42
  39. package/dist/client/mockModel.d.ts.map +0 -1
  40. package/dist/client/mockModel.js +0 -182
  41. package/dist/client/mockModel.js.map +0 -1
  42. package/dist/client/saveInputMessages.d.ts +0 -20
  43. package/dist/client/saveInputMessages.d.ts.map +0 -1
  44. package/dist/client/saveInputMessages.js +0 -58
  45. package/dist/client/saveInputMessages.js.map +0 -1
  46. package/dist/client/search.d.ts +0 -493
  47. package/dist/client/search.d.ts.map +0 -1
  48. package/dist/client/search.js +0 -425
  49. package/dist/client/search.js.map +0 -1
  50. package/dist/client/start.d.ts +0 -84
  51. package/dist/client/start.d.ts.map +0 -1
  52. package/dist/client/start.js +0 -185
  53. package/dist/client/start.js.map +0 -1
  54. package/dist/client/streamText.d.ts +0 -46
  55. package/dist/client/streamText.d.ts.map +0 -1
  56. package/dist/client/streamText.js +0 -117
  57. package/dist/client/streamText.js.map +0 -1
  58. package/dist/client/streaming.d.ts +0 -3778
  59. package/dist/client/streaming.d.ts.map +0 -1
  60. package/dist/client/streaming.js +0 -314
  61. package/dist/client/streaming.js.map +0 -1
  62. package/dist/client/threads.d.ts +0 -46
  63. package/dist/client/threads.d.ts.map +0 -1
  64. package/dist/client/threads.js +0 -49
  65. package/dist/client/threads.js.map +0 -1
  66. package/dist/client/types.d.ts +0 -461
  67. package/dist/client/types.d.ts.map +0 -1
  68. package/dist/client/types.js +0 -2
  69. package/dist/client/types.js.map +0 -1
  70. package/dist/client/utils.d.ts +0 -4
  71. package/dist/client/utils.d.ts.map +0 -1
  72. package/dist/client/utils.js +0 -21
  73. package/dist/client/utils.js.map +0 -1
  74. package/dist/component/_generated/api.d.ts +0 -48
  75. package/dist/component/_generated/api.d.ts.map +0 -1
  76. package/dist/component/_generated/api.js +0 -31
  77. package/dist/component/_generated/api.js.map +0 -1
  78. package/dist/component/_generated/component.d.ts +0 -3120
  79. package/dist/component/_generated/component.d.ts.map +0 -1
  80. package/dist/component/_generated/component.js +0 -11
  81. package/dist/component/_generated/component.js.map +0 -1
  82. package/dist/component/_generated/dataModel.d.ts +0 -46
  83. package/dist/component/_generated/dataModel.d.ts.map +0 -1
  84. package/dist/component/_generated/dataModel.js +0 -11
  85. package/dist/component/_generated/dataModel.js.map +0 -1
  86. package/dist/component/_generated/server.d.ts +0 -121
  87. package/dist/component/_generated/server.d.ts.map +0 -1
  88. package/dist/component/_generated/server.js +0 -78
  89. package/dist/component/_generated/server.js.map +0 -1
  90. package/dist/component/apiKeys.d.ts +0 -11
  91. package/dist/component/apiKeys.d.ts.map +0 -1
  92. package/dist/component/apiKeys.js +0 -69
  93. package/dist/component/apiKeys.js.map +0 -1
  94. package/dist/component/convex.config.d.ts +0 -3
  95. package/dist/component/convex.config.d.ts.map +0 -1
  96. package/dist/component/convex.config.js +0 -3
  97. package/dist/component/convex.config.js.map +0 -1
  98. package/dist/component/files.d.ts +0 -97
  99. package/dist/component/files.d.ts.map +0 -1
  100. package/dist/component/files.js +0 -190
  101. package/dist/component/files.js.map +0 -1
  102. package/dist/component/messages.d.ts +0 -3851
  103. package/dist/component/messages.d.ts.map +0 -1
  104. package/dist/component/messages.js +0 -757
  105. package/dist/component/messages.js.map +0 -1
  106. package/dist/component/schema.d.ts +0 -8026
  107. package/dist/component/schema.d.ts.map +0 -1
  108. package/dist/component/schema.js +0 -147
  109. package/dist/component/schema.js.map +0 -1
  110. package/dist/component/streams.d.ts +0 -128
  111. package/dist/component/streams.d.ts.map +0 -1
  112. package/dist/component/streams.js +0 -413
  113. package/dist/component/streams.js.map +0 -1
  114. package/dist/component/threads.d.ts +0 -115
  115. package/dist/component/threads.d.ts.map +0 -1
  116. package/dist/component/threads.js +0 -208
  117. package/dist/component/threads.js.map +0 -1
  118. package/dist/component/users.d.ts +0 -52
  119. package/dist/component/users.d.ts.map +0 -1
  120. package/dist/component/users.js +0 -229
  121. package/dist/component/users.js.map +0 -1
  122. package/dist/component/vector/index.d.ts +0 -61
  123. package/dist/component/vector/index.d.ts.map +0 -1
  124. package/dist/component/vector/index.js +0 -146
  125. package/dist/component/vector/index.js.map +0 -1
  126. package/dist/component/vector/tables.d.ts +0 -58
  127. package/dist/component/vector/tables.d.ts.map +0 -1
  128. package/dist/component/vector/tables.js +0 -56
  129. package/dist/component/vector/tables.js.map +0 -1
  130. package/dist/deltas.d.ts +0 -43
  131. package/dist/deltas.d.ts.map +0 -1
  132. package/dist/deltas.js +0 -446
  133. package/dist/deltas.js.map +0 -1
  134. package/dist/mapping.d.ts +0 -72
  135. package/dist/mapping.d.ts.map +0 -1
  136. package/dist/mapping.js +0 -677
  137. package/dist/mapping.js.map +0 -1
  138. package/dist/react/SmoothText.d.ts +0 -5
  139. package/dist/react/SmoothText.d.ts.map +0 -1
  140. package/dist/react/SmoothText.js +0 -6
  141. package/dist/react/SmoothText.js.map +0 -1
  142. package/dist/react/index.d.ts +0 -25
  143. package/dist/react/index.d.ts.map +0 -1
  144. package/dist/react/index.js +0 -70
  145. package/dist/react/index.js.map +0 -1
  146. package/dist/react/optimisticallySendMessage.d.ts +0 -42
  147. package/dist/react/optimisticallySendMessage.d.ts.map +0 -1
  148. package/dist/react/optimisticallySendMessage.js +0 -74
  149. package/dist/react/optimisticallySendMessage.js.map +0 -1
  150. package/dist/react/types.d.ts +0 -12
  151. package/dist/react/types.d.ts.map +0 -1
  152. package/dist/react/types.js +0 -2
  153. package/dist/react/types.js.map +0 -1
  154. package/dist/react/useDeltaStreams.d.ts +0 -10
  155. package/dist/react/useDeltaStreams.d.ts.map +0 -1
  156. package/dist/react/useDeltaStreams.js +0 -106
  157. package/dist/react/useDeltaStreams.js.map +0 -1
  158. package/dist/react/useSmoothText.d.ts +0 -27
  159. package/dist/react/useSmoothText.d.ts.map +0 -1
  160. package/dist/react/useSmoothText.js +0 -68
  161. package/dist/react/useSmoothText.js.map +0 -1
  162. package/dist/react/useStreamingUIMessages.d.ts +0 -22
  163. package/dist/react/useStreamingUIMessages.d.ts.map +0 -1
  164. package/dist/react/useStreamingUIMessages.js +0 -92
  165. package/dist/react/useStreamingUIMessages.js.map +0 -1
  166. package/dist/react/useThreadMessages.d.ts +0 -104
  167. package/dist/react/useThreadMessages.d.ts.map +0 -1
  168. package/dist/react/useThreadMessages.js +0 -148
  169. package/dist/react/useThreadMessages.js.map +0 -1
  170. package/dist/react/useUIMessages.d.ts +0 -96
  171. package/dist/react/useUIMessages.d.ts.map +0 -1
  172. package/dist/react/useUIMessages.js +0 -108
  173. package/dist/react/useUIMessages.js.map +0 -1
  174. package/dist/shared.d.ts +0 -26
  175. package/dist/shared.d.ts.map +0 -1
  176. package/dist/shared.js +0 -67
  177. package/dist/shared.js.map +0 -1
  178. package/dist/validators.d.ts +0 -24516
  179. package/dist/validators.d.ts.map +0 -1
  180. package/dist/validators.js +0 -475
  181. package/dist/validators.js.map +0 -1
@@ -1,757 +0,0 @@
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