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

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (235) hide show
  1. package/MIGRATION.md +153 -0
  2. package/README.md +32 -27
  3. package/dist/UIMessages.d.ts +46 -0
  4. package/dist/UIMessages.d.ts.map +1 -0
  5. package/dist/UIMessages.js +546 -0
  6. package/dist/UIMessages.js.map +1 -0
  7. package/dist/client/createTool.d.ts +126 -27
  8. package/dist/client/createTool.d.ts.map +1 -1
  9. package/dist/client/createTool.js +67 -12
  10. package/dist/client/createTool.js.map +1 -1
  11. package/dist/client/defaultComponent.d.ts +11 -0
  12. package/dist/client/defaultComponent.d.ts.map +1 -0
  13. package/dist/client/defaultComponent.js +7 -0
  14. package/dist/client/defaultComponent.js.map +1 -0
  15. package/dist/client/definePlaygroundAPI.d.ts +1335 -204
  16. package/dist/client/definePlaygroundAPI.d.ts.map +1 -1
  17. package/dist/client/definePlaygroundAPI.js +52 -28
  18. package/dist/client/definePlaygroundAPI.js.map +1 -1
  19. package/dist/client/files.d.ts +20 -7
  20. package/dist/client/files.d.ts.map +1 -1
  21. package/dist/client/files.js +68 -11
  22. package/dist/client/files.js.map +1 -1
  23. package/dist/client/index.d.ts +1116 -978
  24. package/dist/client/index.d.ts.map +1 -1
  25. package/dist/client/index.js +332 -747
  26. package/dist/client/index.js.map +1 -1
  27. package/dist/client/messages.d.ts +461 -0
  28. package/dist/client/messages.d.ts.map +1 -0
  29. package/dist/client/messages.js +106 -0
  30. package/dist/client/messages.js.map +1 -0
  31. package/dist/client/mockModel.d.ts +42 -0
  32. package/dist/client/mockModel.d.ts.map +1 -0
  33. package/dist/client/mockModel.js +182 -0
  34. package/dist/client/mockModel.js.map +1 -0
  35. package/dist/client/saveInputMessages.d.ts +20 -0
  36. package/dist/client/saveInputMessages.d.ts.map +1 -0
  37. package/dist/client/saveInputMessages.js +58 -0
  38. package/dist/client/saveInputMessages.js.map +1 -0
  39. package/dist/client/search.d.ts +350 -39
  40. package/dist/client/search.d.ts.map +1 -1
  41. package/dist/client/search.js +350 -39
  42. package/dist/client/search.js.map +1 -1
  43. package/dist/client/start.d.ts +84 -0
  44. package/dist/client/start.d.ts.map +1 -0
  45. package/dist/client/start.js +185 -0
  46. package/dist/client/start.js.map +1 -0
  47. package/dist/client/streamText.d.ts +46 -0
  48. package/dist/client/streamText.d.ts.map +1 -0
  49. package/dist/client/streamText.js +117 -0
  50. package/dist/client/streamText.js.map +1 -0
  51. package/dist/client/streaming.d.ts +3716 -32
  52. package/dist/client/streaming.d.ts.map +1 -1
  53. package/dist/client/streaming.js +161 -59
  54. package/dist/client/streaming.js.map +1 -1
  55. package/dist/client/threads.d.ts +46 -0
  56. package/dist/client/threads.d.ts.map +1 -0
  57. package/dist/client/threads.js +49 -0
  58. package/dist/client/threads.js.map +1 -0
  59. package/dist/client/types.d.ts +266 -128
  60. package/dist/client/types.d.ts.map +1 -1
  61. package/dist/client/utils.d.ts +4 -0
  62. package/dist/client/utils.d.ts.map +1 -0
  63. package/dist/client/utils.js +21 -0
  64. package/dist/client/utils.js.map +1 -0
  65. package/dist/component/_generated/api.d.ts +24 -2178
  66. package/dist/component/_generated/api.d.ts.map +1 -1
  67. package/dist/component/_generated/api.js +10 -1
  68. package/dist/component/_generated/api.js.map +1 -1
  69. package/dist/component/_generated/component.d.ts +3120 -0
  70. package/dist/component/_generated/component.d.ts.map +1 -0
  71. package/dist/component/_generated/component.js +11 -0
  72. package/dist/component/_generated/component.js.map +1 -0
  73. package/dist/component/_generated/dataModel.d.ts +4 -18
  74. package/dist/component/_generated/dataModel.d.ts.map +1 -0
  75. package/dist/component/_generated/dataModel.js +11 -0
  76. package/dist/component/_generated/dataModel.js.map +1 -0
  77. package/dist/component/_generated/server.d.ts +10 -38
  78. package/dist/component/_generated/server.d.ts.map +1 -1
  79. package/dist/component/_generated/server.js +9 -5
  80. package/dist/component/_generated/server.js.map +1 -1
  81. package/dist/component/files.d.ts +16 -10
  82. package/dist/component/files.d.ts.map +1 -1
  83. package/dist/component/files.js +10 -2
  84. package/dist/component/files.js.map +1 -1
  85. package/dist/component/messages.d.ts +2578 -366
  86. package/dist/component/messages.d.ts.map +1 -1
  87. package/dist/component/messages.js +397 -154
  88. package/dist/component/messages.js.map +1 -1
  89. package/dist/component/schema.d.ts +5697 -3584
  90. package/dist/component/schema.d.ts.map +1 -1
  91. package/dist/component/schema.js +18 -41
  92. package/dist/component/schema.js.map +1 -1
  93. package/dist/component/streams.d.ts +39 -339
  94. package/dist/component/streams.d.ts.map +1 -1
  95. package/dist/component/streams.js +114 -73
  96. package/dist/component/streams.js.map +1 -1
  97. package/dist/component/threads.d.ts +13 -13
  98. package/dist/component/users.d.ts +7 -7
  99. package/dist/component/vector/index.d.ts +1 -1
  100. package/dist/component/vector/index.d.ts.map +1 -1
  101. package/dist/component/vector/index.js +1 -3
  102. package/dist/component/vector/index.js.map +1 -1
  103. package/dist/deltas.d.ts +43 -0
  104. package/dist/deltas.d.ts.map +1 -0
  105. package/dist/deltas.js +446 -0
  106. package/dist/deltas.js.map +1 -0
  107. package/dist/mapping.d.ts +38 -20
  108. package/dist/mapping.d.ts.map +1 -1
  109. package/dist/mapping.js +365 -97
  110. package/dist/mapping.js.map +1 -1
  111. package/dist/react/SmoothText.d.ts +5 -0
  112. package/dist/react/SmoothText.d.ts.map +1 -0
  113. package/dist/react/SmoothText.js +6 -0
  114. package/dist/react/SmoothText.js.map +1 -0
  115. package/dist/react/index.d.ts +5 -77
  116. package/dist/react/index.d.ts.map +1 -1
  117. package/dist/react/index.js +6 -160
  118. package/dist/react/index.js.map +1 -1
  119. package/dist/react/optimisticallySendMessage.d.ts +36 -3
  120. package/dist/react/optimisticallySendMessage.d.ts.map +1 -1
  121. package/dist/react/optimisticallySendMessage.js +35 -9
  122. package/dist/react/optimisticallySendMessage.js.map +1 -1
  123. package/dist/react/types.d.ts +4 -18
  124. package/dist/react/types.d.ts.map +1 -1
  125. package/dist/react/useDeltaStreams.d.ts +10 -0
  126. package/dist/react/useDeltaStreams.d.ts.map +1 -0
  127. package/dist/react/useDeltaStreams.js +106 -0
  128. package/dist/react/useDeltaStreams.js.map +1 -0
  129. package/dist/react/useSmoothText.d.ts +13 -12
  130. package/dist/react/useSmoothText.d.ts.map +1 -1
  131. package/dist/react/useSmoothText.js +32 -15
  132. package/dist/react/useSmoothText.js.map +1 -1
  133. package/dist/react/useStreamingUIMessages.d.ts +22 -0
  134. package/dist/react/useStreamingUIMessages.d.ts.map +1 -0
  135. package/dist/react/useStreamingUIMessages.js +92 -0
  136. package/dist/react/useStreamingUIMessages.js.map +1 -0
  137. package/dist/react/useThreadMessages.d.ts +104 -0
  138. package/dist/react/useThreadMessages.d.ts.map +1 -0
  139. package/dist/react/useThreadMessages.js +148 -0
  140. package/dist/react/useThreadMessages.js.map +1 -0
  141. package/dist/react/useUIMessages.d.ts +96 -0
  142. package/dist/react/useUIMessages.d.ts.map +1 -0
  143. package/dist/react/useUIMessages.js +108 -0
  144. package/dist/react/useUIMessages.js.map +1 -0
  145. package/dist/shared.d.ts +20 -4
  146. package/dist/shared.d.ts.map +1 -1
  147. package/dist/shared.js +45 -8
  148. package/dist/shared.js.map +1 -1
  149. package/dist/validators.d.ts +22981 -5666
  150. package/dist/validators.d.ts.map +1 -1
  151. package/dist/validators.js +245 -137
  152. package/dist/validators.js.map +1 -1
  153. package/package.json +101 -51
  154. package/src/UIMessages.combineUIMessages.test.ts +239 -0
  155. package/src/UIMessages.test.ts +273 -0
  156. package/src/UIMessages.ts +739 -0
  157. package/src/client/approval.test.ts +350 -0
  158. package/src/client/createTool.ts +291 -76
  159. package/src/client/defaultComponent.ts +17 -0
  160. package/src/client/definePlaygroundAPI.ts +67 -31
  161. package/src/client/files.ts +100 -20
  162. package/src/client/index.test.ts +40 -85
  163. package/src/client/index.ts +638 -1289
  164. package/src/client/messages.ts +237 -0
  165. package/src/client/mockModel.ts +252 -0
  166. package/src/client/saveInputMessages.test.ts +583 -0
  167. package/src/client/saveInputMessages.ts +101 -0
  168. package/src/client/search.test.ts +1207 -0
  169. package/src/client/search.ts +581 -70
  170. package/src/client/start.ts +327 -0
  171. package/src/client/streamText.ts +187 -0
  172. package/src/client/streaming.test.ts +186 -0
  173. package/src/client/streaming.ts +241 -97
  174. package/src/client/threads.ts +83 -0
  175. package/src/client/types.ts +370 -219
  176. package/src/client/utils.ts +27 -0
  177. package/src/component/_generated/api.ts +64 -0
  178. package/src/component/_generated/component.ts +4902 -0
  179. package/src/component/_generated/{server.d.ts → server.ts} +33 -21
  180. package/src/component/files.ts +11 -2
  181. package/src/component/messages.test.ts +195 -51
  182. package/src/component/messages.ts +500 -201
  183. package/src/component/schema.ts +20 -46
  184. package/src/component/setup.test.ts +7 -0
  185. package/src/component/streams.ts +184 -83
  186. package/src/component/users.test.ts +0 -1
  187. package/src/component/vector/index.ts +1 -3
  188. package/src/deltas.test.ts +626 -0
  189. package/src/deltas.ts +569 -0
  190. package/src/fromUIMessages.test.ts +497 -0
  191. package/src/mapping.test.ts +180 -6
  192. package/src/mapping.ts +479 -162
  193. package/src/react/SmoothText.tsx +9 -0
  194. package/src/react/index.ts +10 -230
  195. package/src/react/optimisticallySendMessage.ts +55 -12
  196. package/src/react/types.ts +6 -39
  197. package/src/react/useDeltaStreams.ts +160 -0
  198. package/src/react/useSmoothText.ts +56 -36
  199. package/src/react/useStreamingUIMessages.ts +143 -0
  200. package/src/react/useThreadMessages.ts +262 -0
  201. package/src/react/useUIMessages.test.ts +255 -0
  202. package/src/react/useUIMessages.ts +195 -0
  203. package/src/shared.ts +88 -12
  204. package/src/test.ts +18 -0
  205. package/src/toUIMessages.test.ts +1269 -0
  206. package/src/validators.test.ts +18 -19
  207. package/src/validators.ts +325 -185
  208. package/dist/client/_generated/_ignore.d.ts +0 -1
  209. package/dist/client/_generated/_ignore.d.ts.map +0 -1
  210. package/dist/client/_generated/_ignore.js +0 -3
  211. package/dist/client/_generated/_ignore.js.map +0 -1
  212. package/dist/client/listMessages.d.ts +0 -22
  213. package/dist/client/listMessages.d.ts.map +0 -1
  214. package/dist/client/listMessages.js +0 -25
  215. package/dist/client/listMessages.js.map +0 -1
  216. package/dist/package.json +0 -3
  217. package/dist/react/deltas.d.ts +0 -26
  218. package/dist/react/deltas.d.ts.map +0 -1
  219. package/dist/react/deltas.js +0 -384
  220. package/dist/react/deltas.js.map +0 -1
  221. package/dist/react/toUIMessages.d.ts +0 -15
  222. package/dist/react/toUIMessages.d.ts.map +0 -1
  223. package/dist/react/toUIMessages.js +0 -211
  224. package/dist/react/toUIMessages.js.map +0 -1
  225. package/src/client/listMessages.ts +0 -38
  226. package/src/component/_generated/api.d.ts +0 -2202
  227. package/src/component/_generated/api.js +0 -23
  228. package/src/component/_generated/server.js +0 -90
  229. package/src/node_modules/.vite/vitest/da39a3ee5e6b4b0d3255bfef95601890afd80709/results.json +0 -1
  230. package/src/react/deltas.test.ts +0 -315
  231. package/src/react/deltas.ts +0 -478
  232. package/src/react/toUIMessages.test.ts +0 -420
  233. package/src/react/toUIMessages.ts +0 -253
  234. package/src/vitest.config.ts +0 -7
  235. /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, finishHandler } 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,21 @@ 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()),
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")),
86
97
  };
87
98
  export const addMessages = mutation({
88
99
  args: addMessagesArgs,
89
100
  handler: addMessagesHandler,
90
- returns: v.object({
91
- messages: v.array(vMessageDoc),
92
- }),
101
+ returns: v.object({ messages: v.array(vMessageDoc) }),
93
102
  });
94
103
  async function addMessagesHandler(ctx, args) {
95
104
  let userId = args.userId;
@@ -99,34 +108,49 @@ async function addMessagesHandler(ctx, args) {
99
108
  assert(thread, `Thread ${args.threadId} not found`);
100
109
  userId = thread.userId;
101
110
  }
102
- const { embeddings, failPendingSteps, pending, messages, promptMessageId, ...rest } = args;
103
- const parentMessage = promptMessageId && (await ctx.db.get(promptMessageId));
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));
104
115
  if (failPendingSteps) {
105
116
  assert(args.threadId, "threadId is required to fail pending steps");
106
117
  const pendingMessages = await ctx.db
107
118
  .query("messages")
108
119
  .withIndex("threadId_status_tool_order_stepOrder", (q) => q.eq("threadId", threadId).eq("status", "pending"))
109
- .collect();
120
+ .order("desc")
121
+ .take(100);
110
122
  await Promise.all(pendingMessages
111
- .filter((m) => !parentMessage || m.order === parentMessage.order)
112
- .map((m) => ctx.db.patch(m._id, { status: "failed", error: "Restarting" })));
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
+ }));
113
135
  }
114
136
  let order, stepOrder;
115
137
  let fail = false;
138
+ let error;
116
139
  if (promptMessageId) {
117
- assert(parentMessage, `Parent message ${promptMessageId} not found`);
118
- if (parentMessage.status === "failed") {
140
+ assert(promptMessage, `Parent message ${promptMessageId} not found`);
141
+ if (promptMessage.status === "failed") {
119
142
  fail = true;
143
+ error = promptMessage.error ?? error ?? "The prompt message failed";
120
144
  }
121
- order = parentMessage.order;
145
+ order = promptMessage.order;
122
146
  // Defend against there being existing messages with this parent.
123
147
  const maxMessage = await getMaxMessage(ctx, threadId, order);
124
- stepOrder = maxMessage?.stepOrder ?? parentMessage.stepOrder;
148
+ stepOrder = maxMessage?.stepOrder ?? promptMessage.stepOrder;
125
149
  }
126
150
  else {
127
151
  const maxMessage = await getMaxMessage(ctx, threadId);
128
- order = maxMessage ? maxMessage.order + 1 : 0;
129
- stepOrder = -1;
152
+ order = maxMessage?.order ?? -1;
153
+ stepOrder = maxMessage?.stepOrder ?? -1;
130
154
  }
131
155
  const toReturn = [];
132
156
  if (embeddings) {
@@ -135,99 +159,174 @@ async function addMessagesHandler(ctx, args) {
135
159
  for (let i = 0; i < messages.length; i++) {
136
160
  const message = messages[i];
137
161
  let embeddingId;
138
- if (embeddings && embeddings.vectors[i]) {
162
+ if (embeddings &&
163
+ embeddings.vectors[i] &&
164
+ !fail &&
165
+ message.status !== "failed") {
139
166
  embeddingId = await insertVector(ctx, embeddings.dimension, {
140
167
  vector: embeddings.vectors[i],
141
168
  model: embeddings.model,
142
169
  table: "messages",
143
- userId,
170
+ userId: hideFromUserIdSearch ? undefined : userId,
144
171
  threadId,
145
172
  });
146
173
  }
147
- stepOrder++;
148
- const messageId = await ctx.db.insert("messages", {
174
+ const messageDoc = {
149
175
  ...rest,
150
176
  ...message,
151
177
  embeddingId,
152
178
  parentMessageId: promptMessageId,
153
179
  userId,
154
- order,
155
180
  tool: isTool(message.message),
156
- text: extractText(message.message),
157
- status: fail ? "failed" : pending ? "pending" : "success",
158
- error: fail ? "Parent message failed" : undefined,
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,
159
229
  stepOrder,
160
230
  });
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
231
  if (message.fileIds) {
168
232
  await changeRefcount(ctx, [], message.fileIds);
169
233
  }
170
234
  // TODO: delete the associated stream data for the order/stepOrder
171
235
  toReturn.push((await ctx.db.get(messageId)));
172
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
+ }
173
242
  return { messages: toReturn.map(publicMessage) };
174
243
  }
175
244
  // exported for tests
176
245
  export async function getMaxMessage(ctx, threadId, order) {
177
- return orderedMessagesStream(ctx, threadId, "desc", order).first();
246
+ return orderedMessagesStream(ctx, {
247
+ threadId,
248
+ sortOrder: "desc",
249
+ startOrder: order,
250
+ startOrderBound: "eq",
251
+ }).first();
178
252
  }
179
- function orderedMessagesStream(ctx, threadId, sortOrder, order) {
253
+ function orderedMessagesStream(ctx, args) {
180
254
  return mergedStream([true, false].flatMap((tool) => messageStatuses.map((status) => stream(ctx.db, schema)
181
255
  .query("messages")
182
256
  .withIndex("threadId_status_tool_order_stepOrder", (q) => {
183
257
  const qq = q
184
- .eq("threadId", threadId)
258
+ .eq("threadId", args.threadId)
185
259
  .eq("status", status)
186
260
  .eq("tool", tool);
187
- if (order) {
188
- return qq.eq("order", order);
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
+ }
189
268
  }
190
269
  return qq;
191
270
  })
192
- .order(sortOrder))), ["order", "stepOrder"]);
271
+ .order(args.sortOrder))), ["order", "stepOrder"]);
193
272
  }
194
- export const rollbackMessage = mutation({
273
+ export const finalizeMessage = mutation({
195
274
  args: {
196
275
  messageId: v.id("messages"),
197
- error: v.optional(v.string()),
276
+ result: v.union(v.object({ status: v.literal("success") }), v.object({ status: v.literal("failed"), error: v.string() })),
198
277
  },
199
278
  returns: v.null(),
200
- handler: async (ctx, { messageId, error }) => {
279
+ handler: async (ctx, { messageId, result }) => {
201
280
  const message = await ctx.db.get(messageId);
202
281
  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 });
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;
207
300
  }
208
301
  }
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"),
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
+ }
218
315
  },
219
- returns: v.null(),
220
- handler: commitMessageHandler,
221
316
  });
222
317
  export const updateMessage = mutation({
223
318
  args: {
224
319
  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
- }),
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
+ ]))),
231
330
  },
232
331
  returns: vMessageDoc,
233
332
  handler: async (ctx, args) => {
@@ -236,83 +335,185 @@ export const updateMessage = mutation({
236
335
  if (args.patch.fileIds) {
237
336
  await changeRefcount(ctx, message.fileIds ?? [], args.patch.fileIds);
238
337
  }
239
- const patch = {
240
- ...args.patch,
241
- };
338
+ const patch = { ...args.patch };
242
339
  if (args.patch.message !== undefined) {
243
340
  patch.message = args.patch.message;
244
341
  patch.tool = isTool(args.patch.message);
245
342
  patch.text = extractText(args.patch.message);
246
343
  }
344
+ if (args.patch.status === "failed") {
345
+ if (message.embeddingId) {
346
+ await ctx.db.delete(message.embeddingId);
347
+ }
348
+ patch.embeddingId = undefined;
349
+ }
247
350
  await ctx.db.patch(args.messageId, patch);
248
351
  return publicMessage((await ctx.db.get(args.messageId)));
249
352
  },
250
353
  });
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({
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({
267
370
  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")),
371
+ ...cloneMessageArgs,
372
+ paginationOpts: paginationOptsValidator,
275
373
  },
276
374
  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)
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)
288
390
  .eq("status", status)
289
- .eq("tool", tool);
290
- if (last) {
291
- return qq.lte("order", last.order);
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);
292
400
  }
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
- });
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);
306
477
  return { ...messages, page: messages.page.map(publicMessage) };
307
478
  },
308
479
  returns: vPaginationResult(vMessageDoc),
309
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
+ }
310
513
  export const getMessagesByIds = query({
311
- args: {
312
- messageIds: v.array(v.id("messages")),
313
- },
514
+ args: { messageIds: v.array(v.id("messages")) },
314
515
  handler: async (ctx, args) => {
315
- return await Promise.all(args.messageIds.map((id) => ctx.db.get(id)));
516
+ return (await Promise.all(args.messageIds.map((id) => ctx.db.get(id)))).map((m) => (m ? publicMessage(m) : null));
316
517
  },
317
518
  returns: v.array(v.union(v.null(), vMessageDoc)),
318
519
  });
@@ -320,10 +521,12 @@ export const searchMessages = action({
320
521
  args: {
321
522
  threadId: v.optional(v.id("threads")),
322
523
  searchAllMessagesForUserId: v.optional(v.string()),
323
- beforeMessageId: v.optional(v.id("messages")),
524
+ targetMessageId: v.optional(v.id("messages")),
324
525
  embedding: v.optional(v.array(v.number())),
325
526
  embeddingModel: v.optional(v.string()),
326
527
  text: v.optional(v.string()),
528
+ textSearch: v.optional(v.boolean()),
529
+ vectorSearch: v.optional(v.boolean()),
327
530
  limit: v.number(),
328
531
  vectorScoreThreshold: v.optional(v.number()),
329
532
  messageRange: v.optional(v.object({ before: v.number(), after: v.number() })),
@@ -333,23 +536,34 @@ export const searchMessages = action({
333
536
  assert(args.searchAllMessagesForUserId || args.threadId, "Specify userId or threadId");
334
537
  const limit = args.limit;
335
538
  let textSearchMessages;
336
- if (args.text) {
539
+ if (args.textSearch) {
337
540
  textSearchMessages = await ctx.runQuery(api.messages.textSearch, {
338
541
  searchAllMessagesForUserId: args.searchAllMessagesForUserId,
339
542
  threadId: args.threadId,
543
+ targetMessageId: args.targetMessageId,
340
544
  text: args.text,
341
545
  limit,
342
- beforeMessageId: args.beforeMessageId,
343
546
  });
344
547
  }
345
- if (args.embedding) {
346
- const dimension = args.embedding.length;
347
- if (!VectorDimensions.includes(dimension)) {
348
- throw new Error(`Unsupported embedding dimension: ${dimension}`);
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
+ }
349
560
  }
350
- const vectors = (await searchVectors(ctx, args.embedding, {
561
+ assert(embedding && model, "Embedding missing");
562
+ const dimension = embedding.length;
563
+ validateVectorDimension(dimension);
564
+ const vectors = (await searchVectors(ctx, embedding, {
351
565
  dimension,
352
- model: args.embeddingModel ?? "unknown",
566
+ model,
353
567
  table: "messages",
354
568
  searchAllMessagesForUserId: args.searchAllMessagesForUserId,
355
569
  threadId: args.threadId,
@@ -372,7 +586,7 @@ export const searchMessages = action({
372
586
  embeddingIds,
373
587
  textSearchMessages: textSearchMessages?.filter((m) => !embeddingIds.includes(m.embeddingId)),
374
588
  messageRange: args.messageRange ?? DEFAULT_MESSAGE_RANGE,
375
- beforeMessageId: args.beforeMessageId,
589
+ beforeMessageId: args.targetMessageId,
376
590
  limit,
377
591
  });
378
592
  return messages;
@@ -414,11 +628,11 @@ export const _fetchSearchMessages = internalQuery({
414
628
  .map(publicMessage);
415
629
  messages.push(...(args.textSearchMessages ?? []));
416
630
  // TODO: prioritize more recent messages
417
- messages.sort((a, b) => a.order - b.order);
631
+ messages = sorted(messages);
418
632
  messages = messages.slice(0, args.limit);
419
633
  // Fetch the surrounding messages
420
634
  if (!threadId) {
421
- return messages.sort((a, b) => a.order - b.order);
635
+ return messages;
422
636
  }
423
637
  const included = {};
424
638
  for (const m of messages) {
@@ -469,7 +683,7 @@ export const _fetchSearchMessages = internalQuery({
469
683
  messages.push(publicMessage(r));
470
684
  }
471
685
  }
472
- return messages.sort((a, b) => a.order - b.order);
686
+ return sorted(messages);
473
687
  },
474
688
  });
475
689
  // returns ranges of messages in order of text search relevance,
@@ -478,21 +692,24 @@ export const textSearch = query({
478
692
  args: {
479
693
  threadId: v.optional(v.id("threads")),
480
694
  searchAllMessagesForUserId: v.optional(v.string()),
481
- text: v.string(),
695
+ text: v.optional(v.string()),
696
+ targetMessageId: v.optional(v.id("messages")),
482
697
  limit: v.number(),
483
- beforeMessageId: v.optional(v.id("messages")),
484
698
  },
485
699
  handler: async (ctx, args) => {
486
700
  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;
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
+ }
489
708
  const messages = await ctx.db
490
709
  .query("messages")
491
710
  .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))
711
+ ? q.search("text", text).eq("userId", args.searchAllMessagesForUserId)
712
+ : q.search("text", text).eq("threadId", args.threadId))
496
713
  // Just in case tool messages slip through
497
714
  .filter((q) => {
498
715
  const qq = q.eq(q.field("tool"), false);
@@ -503,12 +720,38 @@ export const textSearch = query({
503
720
  })
504
721
  .take(args.limit);
505
722
  return messages
506
- .filter((m) => !beforeMessage ||
507
- m.order < beforeMessage.order ||
508
- (m.order === beforeMessage.order &&
509
- m.stepOrder < beforeMessage.stepOrder))
723
+ .filter((m) => !targetMessage ||
724
+ m.order < targetMessage.order ||
725
+ (m.order === targetMessage.order &&
726
+ m.stepOrder < targetMessage.stepOrder))
510
727
  .map(publicMessage);
511
728
  },
512
729
  returns: v.array(vMessageDoc),
513
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
+ });
514
757
  //# sourceMappingURL=messages.js.map