@convex-dev/agent 0.5.0-alpha.1 → 0.6.0-alpha.0

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (233) hide show
  1. package/README.md +32 -27
  2. package/dist/UIMessages.d.ts +46 -0
  3. package/dist/UIMessages.d.ts.map +1 -0
  4. package/dist/UIMessages.js +546 -0
  5. package/dist/UIMessages.js.map +1 -0
  6. package/dist/client/createTool.d.ts +129 -27
  7. package/dist/client/createTool.d.ts.map +1 -1
  8. package/dist/client/createTool.js +66 -12
  9. package/dist/client/createTool.js.map +1 -1
  10. package/dist/client/defaultComponent.d.ts +11 -0
  11. package/dist/client/defaultComponent.d.ts.map +1 -0
  12. package/dist/client/defaultComponent.js +7 -0
  13. package/dist/client/defaultComponent.js.map +1 -0
  14. package/dist/client/definePlaygroundAPI.d.ts +1323 -192
  15. package/dist/client/definePlaygroundAPI.d.ts.map +1 -1
  16. package/dist/client/definePlaygroundAPI.js +52 -28
  17. package/dist/client/definePlaygroundAPI.js.map +1 -1
  18. package/dist/client/files.d.ts +20 -7
  19. package/dist/client/files.d.ts.map +1 -1
  20. package/dist/client/files.js +68 -11
  21. package/dist/client/files.js.map +1 -1
  22. package/dist/client/index.d.ts +1056 -965
  23. package/dist/client/index.d.ts.map +1 -1
  24. package/dist/client/index.js +242 -748
  25. package/dist/client/index.js.map +1 -1
  26. package/dist/client/messages.d.ts +461 -0
  27. package/dist/client/messages.d.ts.map +1 -0
  28. package/dist/client/messages.js +106 -0
  29. package/dist/client/messages.js.map +1 -0
  30. package/dist/client/mockModel.d.ts +42 -0
  31. package/dist/client/mockModel.d.ts.map +1 -0
  32. package/dist/client/mockModel.js +175 -0
  33. package/dist/client/mockModel.js.map +1 -0
  34. package/dist/client/saveInputMessages.d.ts +20 -0
  35. package/dist/client/saveInputMessages.d.ts.map +1 -0
  36. package/dist/client/saveInputMessages.js +58 -0
  37. package/dist/client/saveInputMessages.js.map +1 -0
  38. package/dist/client/search.d.ts +346 -35
  39. package/dist/client/search.d.ts.map +1 -1
  40. package/dist/client/search.js +350 -39
  41. package/dist/client/search.js.map +1 -1
  42. package/dist/client/start.d.ts +84 -0
  43. package/dist/client/start.d.ts.map +1 -0
  44. package/dist/client/start.js +171 -0
  45. package/dist/client/start.js.map +1 -0
  46. package/dist/client/streamText.d.ts +46 -0
  47. package/dist/client/streamText.d.ts.map +1 -0
  48. package/dist/client/streamText.js +93 -0
  49. package/dist/client/streamText.js.map +1 -0
  50. package/dist/client/streaming.d.ts +3705 -32
  51. package/dist/client/streaming.d.ts.map +1 -1
  52. package/dist/client/streaming.js +141 -59
  53. package/dist/client/streaming.js.map +1 -1
  54. package/dist/client/threads.d.ts +46 -0
  55. package/dist/client/threads.d.ts.map +1 -0
  56. package/dist/client/threads.js +49 -0
  57. package/dist/client/threads.js.map +1 -0
  58. package/dist/client/types.d.ts +265 -128
  59. package/dist/client/types.d.ts.map +1 -1
  60. package/dist/client/utils.d.ts +4 -0
  61. package/dist/client/utils.d.ts.map +1 -0
  62. package/dist/client/utils.js +21 -0
  63. package/dist/client/utils.js.map +1 -0
  64. package/dist/component/_generated/api.d.ts +24 -2178
  65. package/dist/component/_generated/api.d.ts.map +1 -1
  66. package/dist/component/_generated/api.js +10 -1
  67. package/dist/component/_generated/api.js.map +1 -1
  68. package/dist/component/_generated/component.d.ts +3119 -0
  69. package/dist/component/_generated/component.d.ts.map +1 -0
  70. package/dist/component/_generated/component.js +11 -0
  71. package/dist/component/_generated/component.js.map +1 -0
  72. package/dist/component/_generated/dataModel.d.ts +4 -18
  73. package/dist/component/_generated/dataModel.d.ts.map +1 -0
  74. package/dist/component/_generated/dataModel.js +11 -0
  75. package/dist/component/_generated/dataModel.js.map +1 -0
  76. package/dist/component/_generated/server.d.ts +10 -38
  77. package/dist/component/_generated/server.d.ts.map +1 -1
  78. package/dist/component/_generated/server.js +9 -5
  79. package/dist/component/_generated/server.js.map +1 -1
  80. package/dist/component/files.d.ts +16 -10
  81. package/dist/component/files.d.ts.map +1 -1
  82. package/dist/component/files.js +10 -2
  83. package/dist/component/files.js.map +1 -1
  84. package/dist/component/messages.d.ts +2553 -342
  85. package/dist/component/messages.d.ts.map +1 -1
  86. package/dist/component/messages.js +387 -154
  87. package/dist/component/messages.js.map +1 -1
  88. package/dist/component/schema.d.ts +5697 -3584
  89. package/dist/component/schema.d.ts.map +1 -1
  90. package/dist/component/schema.js +18 -41
  91. package/dist/component/schema.js.map +1 -1
  92. package/dist/component/streams.d.ts +35 -335
  93. package/dist/component/streams.d.ts.map +1 -1
  94. package/dist/component/streams.js +114 -73
  95. package/dist/component/streams.js.map +1 -1
  96. package/dist/component/threads.d.ts +16 -16
  97. package/dist/component/users.d.ts +4 -4
  98. package/dist/component/vector/index.d.ts +1 -1
  99. package/dist/component/vector/index.d.ts.map +1 -1
  100. package/dist/component/vector/index.js +1 -3
  101. package/dist/component/vector/index.js.map +1 -1
  102. package/dist/deltas.d.ts +43 -0
  103. package/dist/deltas.d.ts.map +1 -0
  104. package/dist/deltas.js +447 -0
  105. package/dist/deltas.js.map +1 -0
  106. package/dist/mapping.d.ts +20 -20
  107. package/dist/mapping.d.ts.map +1 -1
  108. package/dist/mapping.js +313 -96
  109. package/dist/mapping.js.map +1 -1
  110. package/dist/react/SmoothText.d.ts +5 -0
  111. package/dist/react/SmoothText.d.ts.map +1 -0
  112. package/dist/react/SmoothText.js +6 -0
  113. package/dist/react/SmoothText.js.map +1 -0
  114. package/dist/react/index.d.ts +5 -77
  115. package/dist/react/index.d.ts.map +1 -1
  116. package/dist/react/index.js +6 -160
  117. package/dist/react/index.js.map +1 -1
  118. package/dist/react/optimisticallySendMessage.d.ts +36 -3
  119. package/dist/react/optimisticallySendMessage.d.ts.map +1 -1
  120. package/dist/react/optimisticallySendMessage.js +35 -9
  121. package/dist/react/optimisticallySendMessage.js.map +1 -1
  122. package/dist/react/types.d.ts +4 -18
  123. package/dist/react/types.d.ts.map +1 -1
  124. package/dist/react/useDeltaStreams.d.ts +10 -0
  125. package/dist/react/useDeltaStreams.d.ts.map +1 -0
  126. package/dist/react/useDeltaStreams.js +101 -0
  127. package/dist/react/useDeltaStreams.js.map +1 -0
  128. package/dist/react/useSmoothText.d.ts +13 -12
  129. package/dist/react/useSmoothText.d.ts.map +1 -1
  130. package/dist/react/useSmoothText.js +32 -15
  131. package/dist/react/useSmoothText.js.map +1 -1
  132. package/dist/react/useStreamingUIMessages.d.ts +22 -0
  133. package/dist/react/useStreamingUIMessages.d.ts.map +1 -0
  134. package/dist/react/useStreamingUIMessages.js +92 -0
  135. package/dist/react/useStreamingUIMessages.js.map +1 -0
  136. package/dist/react/useThreadMessages.d.ts +104 -0
  137. package/dist/react/useThreadMessages.d.ts.map +1 -0
  138. package/dist/react/useThreadMessages.js +148 -0
  139. package/dist/react/useThreadMessages.js.map +1 -0
  140. package/dist/react/useUIMessages.d.ts +96 -0
  141. package/dist/react/useUIMessages.d.ts.map +1 -0
  142. package/dist/react/useUIMessages.js +108 -0
  143. package/dist/react/useUIMessages.js.map +1 -0
  144. package/dist/shared.d.ts +20 -4
  145. package/dist/shared.d.ts.map +1 -1
  146. package/dist/shared.js +45 -8
  147. package/dist/shared.js.map +1 -1
  148. package/dist/validators.d.ts +22981 -5666
  149. package/dist/validators.d.ts.map +1 -1
  150. package/dist/validators.js +245 -137
  151. package/dist/validators.js.map +1 -1
  152. package/package.json +98 -50
  153. package/src/UIMessages.combineUIMessages.test.ts +239 -0
  154. package/src/UIMessages.test.ts +273 -0
  155. package/src/UIMessages.ts +739 -0
  156. package/src/client/createTool.ts +293 -76
  157. package/src/client/defaultComponent.ts +17 -0
  158. package/src/client/definePlaygroundAPI.ts +67 -31
  159. package/src/client/files.ts +100 -20
  160. package/src/client/index.test.ts +40 -85
  161. package/src/client/index.ts +520 -1290
  162. package/src/client/messages.ts +237 -0
  163. package/src/client/mockModel.ts +245 -0
  164. package/src/client/saveInputMessages.test.ts +583 -0
  165. package/src/client/saveInputMessages.ts +101 -0
  166. package/src/client/search.test.ts +1207 -0
  167. package/src/client/search.ts +577 -70
  168. package/src/client/start.ts +310 -0
  169. package/src/client/streamText.ts +163 -0
  170. package/src/client/streaming.test.ts +186 -0
  171. package/src/client/streaming.ts +219 -97
  172. package/src/client/threads.ts +83 -0
  173. package/src/client/types.ts +368 -219
  174. package/src/client/utils.ts +27 -0
  175. package/src/component/_generated/api.ts +64 -0
  176. package/src/component/_generated/component.ts +4913 -0
  177. package/src/component/_generated/{server.d.ts → server.ts} +33 -21
  178. package/src/component/files.ts +11 -2
  179. package/src/component/messages.test.ts +195 -51
  180. package/src/component/messages.ts +490 -201
  181. package/src/component/schema.ts +20 -46
  182. package/src/component/setup.test.ts +7 -0
  183. package/src/component/streams.ts +184 -83
  184. package/src/component/users.test.ts +0 -1
  185. package/src/component/vector/index.ts +1 -3
  186. package/src/deltas.test.ts +626 -0
  187. package/src/deltas.ts +570 -0
  188. package/src/fromUIMessages.test.ts +497 -0
  189. package/src/mapping.test.ts +103 -6
  190. package/src/mapping.ts +422 -161
  191. package/src/react/SmoothText.tsx +9 -0
  192. package/src/react/index.ts +10 -230
  193. package/src/react/optimisticallySendMessage.ts +55 -12
  194. package/src/react/types.ts +6 -39
  195. package/src/react/useDeltaStreams.ts +154 -0
  196. package/src/react/useSmoothText.ts +56 -36
  197. package/src/react/useStreamingUIMessages.ts +143 -0
  198. package/src/react/useThreadMessages.ts +262 -0
  199. package/src/react/useUIMessages.test.ts +255 -0
  200. package/src/react/useUIMessages.ts +195 -0
  201. package/src/shared.ts +88 -12
  202. package/src/test.ts +18 -0
  203. package/src/toUIMessages.test.ts +1269 -0
  204. package/src/validators.test.ts +18 -19
  205. package/src/validators.ts +325 -185
  206. package/dist/client/_generated/_ignore.d.ts +0 -1
  207. package/dist/client/_generated/_ignore.d.ts.map +0 -1
  208. package/dist/client/_generated/_ignore.js +0 -3
  209. package/dist/client/_generated/_ignore.js.map +0 -1
  210. package/dist/client/listMessages.d.ts +0 -22
  211. package/dist/client/listMessages.d.ts.map +0 -1
  212. package/dist/client/listMessages.js +0 -25
  213. package/dist/client/listMessages.js.map +0 -1
  214. package/dist/package.json +0 -3
  215. package/dist/react/deltas.d.ts +0 -26
  216. package/dist/react/deltas.d.ts.map +0 -1
  217. package/dist/react/deltas.js +0 -384
  218. package/dist/react/deltas.js.map +0 -1
  219. package/dist/react/toUIMessages.d.ts +0 -15
  220. package/dist/react/toUIMessages.d.ts.map +0 -1
  221. package/dist/react/toUIMessages.js +0 -211
  222. package/dist/react/toUIMessages.js.map +0 -1
  223. package/src/client/listMessages.ts +0 -38
  224. package/src/component/_generated/api.d.ts +0 -2202
  225. package/src/component/_generated/api.js +0 -23
  226. package/src/component/_generated/server.js +0 -90
  227. package/src/node_modules/.vite/vitest/da39a3ee5e6b4b0d3255bfef95601890afd80709/results.json +0 -1
  228. package/src/react/deltas.test.ts +0 -315
  229. package/src/react/deltas.ts +0 -478
  230. package/src/react/toUIMessages.test.ts +0 -420
  231. package/src/react/toUIMessages.ts +0 -253
  232. package/src/vitest.config.ts +0 -7
  233. /package/src/component/_generated/{dataModel.d.ts → dataModel.ts} +0 -0
@@ -1,15 +1,16 @@
1
- import { assert, omit } from "convex-helpers";
1
+ import { assert, omit, pick } from "convex-helpers";
2
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, } from "../shared.js";
5
- import { vMessageEmbeddings, vMessageStatus, vMessageWithMetadataInternal, vPaginationResult, } from "../validators.js";
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
6
  import { api, internal } from "./_generated/api.js";
7
- import { action, internalQuery, mutation, query, } from "./_generated/server.js";
8
- import { schema, v, vMessageDoc } from "./schema.js";
9
- import { getThread as _getThread, listThreadsByUserId as _listThreadsByUserId, updateThread as _updateThread, } from "./threads.js";
7
+ import { action, internalMutation, internalQuery, mutation, query, } from "./_generated/server.js";
8
+ import { schema, v } from "./schema.js";
10
9
  import { insertVector, searchVectors } from "./vector/index.js";
11
- import { VectorDimensions, vVectorId, } from "./vector/tables.js";
10
+ import { validateVectorDimension, vVectorId, } from "./vector/tables.js";
12
11
  import { changeRefcount } from "./files.js";
12
+ import { getStreamingMessagesWithMetadata } from "./streams.js";
13
+ import { partial } from "convex-helpers/validators";
13
14
  function publicMessage(message) {
14
15
  return omit(message, ["parentMessageId", "stepId", "files"]);
15
16
  }
@@ -23,9 +24,7 @@ export async function deleteMessage(ctx, messageDoc) {
23
24
  }
24
25
  }
25
26
  export const deleteByIds = mutation({
26
- args: {
27
- messageIds: v.array(v.id("messages")),
28
- },
27
+ args: { messageIds: v.array(v.id("messages")) },
29
28
  returns: v.array(v.id("messages")),
30
29
  handler: async (ctx, args) => {
31
30
  const deletedMessageIds = await Promise.all(args.messageIds.map(async (id) => {
@@ -54,7 +53,12 @@ export const deleteByOrder = mutation({
54
53
  lastStepOrder: v.optional(v.number()),
55
54
  }),
56
55
  handler: async (ctx, args) => {
57
- const messages = await orderedMessagesStream(ctx, args.threadId, "asc", args.startOrder)
56
+ const messages = await orderedMessagesStream(ctx, {
57
+ threadId: args.threadId,
58
+ sortOrder: "asc",
59
+ startOrder: args.startOrder,
60
+ startOrderBound: "gte",
61
+ })
58
62
  .narrow({
59
63
  lowerBound: args.startStepOrder
60
64
  ? [args.startOrder, args.startStepOrder]
@@ -80,16 +84,18 @@ const addMessagesArgs = {
80
84
  promptMessageId: v.optional(v.id("messages")),
81
85
  agentName: v.optional(v.string()),
82
86
  messages: v.array(vMessageWithMetadataInternal),
83
- embeddings: v.optional(vMessageEmbeddings),
84
- pending: v.optional(v.boolean()),
87
+ embeddings: v.optional(vMessageEmbeddingsWithDimension),
85
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()),
86
94
  };
87
95
  export const addMessages = mutation({
88
96
  args: addMessagesArgs,
89
97
  handler: addMessagesHandler,
90
- returns: v.object({
91
- messages: v.array(vMessageDoc),
92
- }),
98
+ returns: v.object({ messages: v.array(vMessageDoc) }),
93
99
  });
94
100
  async function addMessagesHandler(ctx, args) {
95
101
  let userId = args.userId;
@@ -99,34 +105,47 @@ async function addMessagesHandler(ctx, args) {
99
105
  assert(thread, `Thread ${args.threadId} not found`);
100
106
  userId = thread.userId;
101
107
  }
102
- const { embeddings, failPendingSteps, pending, messages, promptMessageId, ...rest } = args;
103
- const parentMessage = promptMessageId && (await ctx.db.get(promptMessageId));
108
+ const { embeddings, failPendingSteps, messages, promptMessageId, pendingMessageId, hideFromUserIdSearch, ...rest } = args;
109
+ const promptMessage = promptMessageId && (await ctx.db.get(promptMessageId));
104
110
  if (failPendingSteps) {
105
111
  assert(args.threadId, "threadId is required to fail pending steps");
106
112
  const pendingMessages = await ctx.db
107
113
  .query("messages")
108
114
  .withIndex("threadId_status_tool_order_stepOrder", (q) => q.eq("threadId", threadId).eq("status", "pending"))
109
- .collect();
115
+ .order("desc")
116
+ .take(100);
110
117
  await Promise.all(pendingMessages
111
- .filter((m) => !parentMessage || m.order === parentMessage.order)
112
- .map((m) => ctx.db.patch(m._id, { status: "failed", error: "Restarting" })));
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
+ }));
113
130
  }
114
131
  let order, stepOrder;
115
132
  let fail = false;
133
+ let error;
116
134
  if (promptMessageId) {
117
- assert(parentMessage, `Parent message ${promptMessageId} not found`);
118
- if (parentMessage.status === "failed") {
135
+ assert(promptMessage, `Parent message ${promptMessageId} not found`);
136
+ if (promptMessage.status === "failed") {
119
137
  fail = true;
138
+ error = promptMessage.error ?? error ?? "The prompt message failed";
120
139
  }
121
- order = parentMessage.order;
140
+ order = promptMessage.order;
122
141
  // Defend against there being existing messages with this parent.
123
142
  const maxMessage = await getMaxMessage(ctx, threadId, order);
124
- stepOrder = maxMessage?.stepOrder ?? parentMessage.stepOrder;
143
+ stepOrder = maxMessage?.stepOrder ?? promptMessage.stepOrder;
125
144
  }
126
145
  else {
127
146
  const maxMessage = await getMaxMessage(ctx, threadId);
128
- order = maxMessage ? maxMessage.order + 1 : 0;
129
- stepOrder = -1;
147
+ order = maxMessage?.order ?? -1;
148
+ stepOrder = maxMessage?.stepOrder ?? -1;
130
149
  }
131
150
  const toReturn = [];
132
151
  if (embeddings) {
@@ -135,35 +154,75 @@ async function addMessagesHandler(ctx, args) {
135
154
  for (let i = 0; i < messages.length; i++) {
136
155
  const message = messages[i];
137
156
  let embeddingId;
138
- if (embeddings && embeddings.vectors[i]) {
157
+ if (embeddings &&
158
+ embeddings.vectors[i] &&
159
+ !fail &&
160
+ message.status !== "failed") {
139
161
  embeddingId = await insertVector(ctx, embeddings.dimension, {
140
162
  vector: embeddings.vectors[i],
141
163
  model: embeddings.model,
142
164
  table: "messages",
143
- userId,
165
+ userId: hideFromUserIdSearch ? undefined : userId,
144
166
  threadId,
145
167
  });
146
168
  }
147
- stepOrder++;
148
- const messageId = await ctx.db.insert("messages", {
169
+ const messageDoc = {
149
170
  ...rest,
150
171
  ...message,
151
172
  embeddingId,
152
173
  parentMessageId: promptMessageId,
153
174
  userId,
154
- order,
155
175
  tool: isTool(message.message),
156
- text: extractText(message.message),
157
- status: fail ? "failed" : pending ? "pending" : "success",
158
- error: fail ? "Parent message failed" : undefined,
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,
159
224
  stepOrder,
160
225
  });
161
- // Let's just not set the id field and have it set only in explicit cases.
162
- // if (!message.id) {
163
- // await ctx.db.patch(messageId, {
164
- // id: messageId,
165
- // });
166
- // }
167
226
  if (message.fileIds) {
168
227
  await changeRefcount(ctx, [], message.fileIds);
169
228
  }
@@ -174,60 +233,90 @@ async function addMessagesHandler(ctx, args) {
174
233
  }
175
234
  // exported for tests
176
235
  export async function getMaxMessage(ctx, threadId, order) {
177
- return orderedMessagesStream(ctx, threadId, "desc", order).first();
236
+ return orderedMessagesStream(ctx, {
237
+ threadId,
238
+ sortOrder: "desc",
239
+ startOrder: order,
240
+ startOrderBound: "eq",
241
+ }).first();
178
242
  }
179
- function orderedMessagesStream(ctx, threadId, sortOrder, order) {
243
+ function orderedMessagesStream(ctx, args) {
180
244
  return mergedStream([true, false].flatMap((tool) => messageStatuses.map((status) => stream(ctx.db, schema)
181
245
  .query("messages")
182
246
  .withIndex("threadId_status_tool_order_stepOrder", (q) => {
183
247
  const qq = q
184
- .eq("threadId", threadId)
248
+ .eq("threadId", args.threadId)
185
249
  .eq("status", status)
186
250
  .eq("tool", tool);
187
- if (order) {
188
- return qq.eq("order", order);
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
+ }
189
258
  }
190
259
  return qq;
191
260
  })
192
- .order(sortOrder))), ["order", "stepOrder"]);
261
+ .order(args.sortOrder))), ["order", "stepOrder"]);
193
262
  }
194
- export const rollbackMessage = mutation({
263
+ export const finalizeMessage = mutation({
195
264
  args: {
196
265
  messageId: v.id("messages"),
197
- error: v.optional(v.string()),
266
+ result: v.union(v.object({ status: v.literal("success") }), v.object({ status: v.literal("failed"), error: v.string() })),
198
267
  },
199
268
  returns: v.null(),
200
- handler: async (ctx, { messageId, error }) => {
269
+ handler: async (ctx, { messageId, result }) => {
201
270
  const message = await ctx.db.get(messageId);
202
271
  assert(message, `Message ${messageId} not found`);
203
- const messages = await orderedMessagesStream(ctx, message.threadId, "asc", message.order).collect();
204
- for (const m of messages) {
205
- if (m.status === "pending") {
206
- await ctx.db.patch(m._id, { status: "failed", error });
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;
207
290
  }
208
291
  }
209
- await ctx.db.patch(messageId, {
210
- status: "failed",
211
- error: error,
212
- });
213
- },
214
- });
215
- export const commitMessage = mutation({
216
- args: {
217
- messageId: v.id("messages"),
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
+ }
218
305
  },
219
- returns: v.null(),
220
- handler: commitMessageHandler,
221
306
  });
222
307
  export const updateMessage = mutation({
223
308
  args: {
224
309
  messageId: v.id("messages"),
225
- patch: v.object({
226
- message: v.optional(vMessageDoc.fields.message),
227
- fileIds: v.optional(v.array(v.id("files"))),
228
- status: v.optional(vMessageStatus),
229
- error: v.optional(v.string()),
230
- }),
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
+ ]))),
231
320
  },
232
321
  returns: vMessageDoc,
233
322
  handler: async (ctx, args) => {
@@ -236,83 +325,185 @@ export const updateMessage = mutation({
236
325
  if (args.patch.fileIds) {
237
326
  await changeRefcount(ctx, message.fileIds ?? [], args.patch.fileIds);
238
327
  }
239
- const patch = {
240
- ...args.patch,
241
- };
328
+ const patch = { ...args.patch };
242
329
  if (args.patch.message !== undefined) {
243
330
  patch.message = args.patch.message;
244
331
  patch.tool = isTool(args.patch.message);
245
332
  patch.text = extractText(args.patch.message);
246
333
  }
334
+ if (args.patch.status === "failed") {
335
+ if (message.embeddingId) {
336
+ await ctx.db.delete(message.embeddingId);
337
+ }
338
+ patch.embeddingId = undefined;
339
+ }
247
340
  await ctx.db.patch(args.messageId, patch);
248
341
  return publicMessage((await ctx.db.get(args.messageId)));
249
342
  },
250
343
  });
251
- async function commitMessageHandler(ctx, { messageId }) {
252
- const message = await ctx.db.get(messageId);
253
- assert(message, `Message ${messageId} not found`);
254
- const order = message.order;
255
- const messages = await mergedStream([true, false].map((tool) => stream(ctx.db, schema)
256
- .query("messages")
257
- .withIndex("threadId_status_tool_order_stepOrder", (q) => q
258
- .eq("threadId", message.threadId)
259
- .eq("status", "pending")
260
- .eq("tool", tool)
261
- .eq("order", order))), ["order", "stepOrder"]).collect();
262
- for (const message of messages) {
263
- await ctx.db.patch(message._id, { status: "success" });
264
- }
265
- }
266
- export const listMessagesByThreadId = query({
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({
267
360
  args: {
268
- threadId: v.id("threads"),
269
- excludeToolMessages: v.optional(v.boolean()),
270
- /** What order to sort the messages in. To get the latest, use "desc". */
271
- order: v.union(v.literal("asc"), v.literal("desc")),
272
- paginationOpts: v.optional(paginationOptsValidator),
273
- statuses: v.optional(v.array(vMessageStatus)),
274
- upToAndIncludingMessageId: v.optional(v.id("messages")),
361
+ ...cloneMessageArgs,
362
+ paginationOpts: paginationOptsValidator,
275
363
  },
276
364
  handler: async (ctx, args) => {
277
- const statuses = args.statuses ?? vMessageStatus.members.map((m) => m.value);
278
- const last = args.upToAndIncludingMessageId &&
279
- (await ctx.db.get(args.upToAndIncludingMessageId));
280
- assert(!last || last.threadId === args.threadId, "upToAndIncludingMessageId must be a message in the thread");
281
- const toolOptions = args.excludeToolMessages ? [false] : [true, false];
282
- const order = args.order ?? "desc";
283
- const streams = toolOptions.flatMap((tool) => statuses.map((status) => stream(ctx.db, schema)
284
- .query("messages")
285
- .withIndex("threadId_status_tool_order_stepOrder", (q) => {
286
- const qq = q
287
- .eq("threadId", args.threadId)
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)
288
380
  .eq("status", status)
289
- .eq("tool", tool);
290
- if (last) {
291
- return qq.lte("order", last.order);
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);
292
390
  }
293
- return qq;
294
- })
295
- .order(order)
296
- .filterWith(
297
- // We allow all messages on the same order.
298
- async (m) => !last || m.order < last.order || m.order === last.order)));
299
- const messages = await mergedStream(streams, [
300
- "order",
301
- "stepOrder",
302
- ]).paginate(args.paginationOpts ?? {
303
- numItems: DEFAULT_RECENT_MESSAGES,
304
- cursor: null,
305
- });
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);
306
467
  return { ...messages, page: messages.page.map(publicMessage) };
307
468
  },
308
469
  returns: vPaginationResult(vMessageDoc),
309
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
+ }
310
503
  export const getMessagesByIds = query({
311
- args: {
312
- messageIds: v.array(v.id("messages")),
313
- },
504
+ args: { messageIds: v.array(v.id("messages")) },
314
505
  handler: async (ctx, args) => {
315
- return await Promise.all(args.messageIds.map((id) => ctx.db.get(id)));
506
+ return (await Promise.all(args.messageIds.map((id) => ctx.db.get(id)))).map((m) => (m ? publicMessage(m) : null));
316
507
  },
317
508
  returns: v.array(v.union(v.null(), vMessageDoc)),
318
509
  });
@@ -320,10 +511,12 @@ export const searchMessages = action({
320
511
  args: {
321
512
  threadId: v.optional(v.id("threads")),
322
513
  searchAllMessagesForUserId: v.optional(v.string()),
323
- beforeMessageId: v.optional(v.id("messages")),
514
+ targetMessageId: v.optional(v.id("messages")),
324
515
  embedding: v.optional(v.array(v.number())),
325
516
  embeddingModel: v.optional(v.string()),
326
517
  text: v.optional(v.string()),
518
+ textSearch: v.optional(v.boolean()),
519
+ vectorSearch: v.optional(v.boolean()),
327
520
  limit: v.number(),
328
521
  vectorScoreThreshold: v.optional(v.number()),
329
522
  messageRange: v.optional(v.object({ before: v.number(), after: v.number() })),
@@ -333,23 +526,34 @@ export const searchMessages = action({
333
526
  assert(args.searchAllMessagesForUserId || args.threadId, "Specify userId or threadId");
334
527
  const limit = args.limit;
335
528
  let textSearchMessages;
336
- if (args.text) {
529
+ if (args.textSearch) {
337
530
  textSearchMessages = await ctx.runQuery(api.messages.textSearch, {
338
531
  searchAllMessagesForUserId: args.searchAllMessagesForUserId,
339
532
  threadId: args.threadId,
533
+ targetMessageId: args.targetMessageId,
340
534
  text: args.text,
341
535
  limit,
342
- beforeMessageId: args.beforeMessageId,
343
536
  });
344
537
  }
345
- if (args.embedding) {
346
- const dimension = args.embedding.length;
347
- if (!VectorDimensions.includes(dimension)) {
348
- throw new Error(`Unsupported embedding dimension: ${dimension}`);
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
+ }
349
550
  }
350
- const vectors = (await searchVectors(ctx, args.embedding, {
551
+ assert(embedding && model, "Embedding missing");
552
+ const dimension = embedding.length;
553
+ validateVectorDimension(dimension);
554
+ const vectors = (await searchVectors(ctx, embedding, {
351
555
  dimension,
352
- model: args.embeddingModel ?? "unknown",
556
+ model,
353
557
  table: "messages",
354
558
  searchAllMessagesForUserId: args.searchAllMessagesForUserId,
355
559
  threadId: args.threadId,
@@ -372,7 +576,7 @@ export const searchMessages = action({
372
576
  embeddingIds,
373
577
  textSearchMessages: textSearchMessages?.filter((m) => !embeddingIds.includes(m.embeddingId)),
374
578
  messageRange: args.messageRange ?? DEFAULT_MESSAGE_RANGE,
375
- beforeMessageId: args.beforeMessageId,
579
+ beforeMessageId: args.targetMessageId,
376
580
  limit,
377
581
  });
378
582
  return messages;
@@ -414,11 +618,11 @@ export const _fetchSearchMessages = internalQuery({
414
618
  .map(publicMessage);
415
619
  messages.push(...(args.textSearchMessages ?? []));
416
620
  // TODO: prioritize more recent messages
417
- messages.sort((a, b) => a.order - b.order);
621
+ messages = sorted(messages);
418
622
  messages = messages.slice(0, args.limit);
419
623
  // Fetch the surrounding messages
420
624
  if (!threadId) {
421
- return messages.sort((a, b) => a.order - b.order);
625
+ return messages;
422
626
  }
423
627
  const included = {};
424
628
  for (const m of messages) {
@@ -469,7 +673,7 @@ export const _fetchSearchMessages = internalQuery({
469
673
  messages.push(publicMessage(r));
470
674
  }
471
675
  }
472
- return messages.sort((a, b) => a.order - b.order);
676
+ return sorted(messages);
473
677
  },
474
678
  });
475
679
  // returns ranges of messages in order of text search relevance,
@@ -478,21 +682,24 @@ export const textSearch = query({
478
682
  args: {
479
683
  threadId: v.optional(v.id("threads")),
480
684
  searchAllMessagesForUserId: v.optional(v.string()),
481
- text: v.string(),
685
+ text: v.optional(v.string()),
686
+ targetMessageId: v.optional(v.id("messages")),
482
687
  limit: v.number(),
483
- beforeMessageId: v.optional(v.id("messages")),
484
688
  },
485
689
  handler: async (ctx, args) => {
486
690
  assert(args.searchAllMessagesForUserId || args.threadId, "Specify userId or threadId");
487
- const beforeMessage = args.beforeMessageId && (await ctx.db.get(args.beforeMessageId));
488
- const order = beforeMessage?.order;
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
+ }
489
698
  const messages = await ctx.db
490
699
  .query("messages")
491
700
  .withSearchIndex("text_search", (q) => args.searchAllMessagesForUserId
492
- ? q
493
- .search("text", args.text)
494
- .eq("userId", args.searchAllMessagesForUserId)
495
- : q.search("text", args.text).eq("threadId", args.threadId))
701
+ ? q.search("text", text).eq("userId", args.searchAllMessagesForUserId)
702
+ : q.search("text", text).eq("threadId", args.threadId))
496
703
  // Just in case tool messages slip through
497
704
  .filter((q) => {
498
705
  const qq = q.eq(q.field("tool"), false);
@@ -503,12 +710,38 @@ export const textSearch = query({
503
710
  })
504
711
  .take(args.limit);
505
712
  return messages
506
- .filter((m) => !beforeMessage ||
507
- m.order < beforeMessage.order ||
508
- (m.order === beforeMessage.order &&
509
- m.stepOrder < beforeMessage.stepOrder))
713
+ .filter((m) => !targetMessage ||
714
+ m.order < targetMessage.order ||
715
+ (m.order === targetMessage.order &&
716
+ m.stepOrder < targetMessage.stepOrder))
510
717
  .map(publicMessage);
511
718
  },
512
719
  returns: v.array(vMessageDoc),
513
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
+ });
514
747
  //# sourceMappingURL=messages.js.map