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