@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,44 +1,46 @@
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";
3
+ import {
4
+ paginationOptsValidator,
5
+ type WithoutSystemFields,
6
+ } from "convex/server";
4
7
  import type { ObjectType } from "convex/values";
5
8
  import {
6
9
  DEFAULT_MESSAGE_RANGE,
7
10
  DEFAULT_RECENT_MESSAGES,
8
11
  extractText,
9
12
  isTool,
13
+ sorted,
10
14
  } from "../shared.js";
11
15
  import {
12
- vMessageEmbeddings,
16
+ vMessageDoc,
17
+ vMessageEmbeddingsWithDimension,
13
18
  vMessageStatus,
14
19
  vMessageWithMetadataInternal,
15
20
  vPaginationResult,
21
+ type MessageDoc,
16
22
  } from "../validators.js";
17
23
  import { api, internal } from "./_generated/api.js";
18
24
  import type { Doc, Id } from "./_generated/dataModel.js";
19
25
  import {
20
26
  action,
27
+ internalMutation,
21
28
  internalQuery,
22
29
  mutation,
23
30
  type MutationCtx,
24
31
  query,
25
32
  type QueryCtx,
26
33
  } from "./_generated/server.js";
27
- import type { MessageDoc } from "./schema.js";
28
- import { schema, v, vMessageDoc } from "./schema.js";
29
- import {
30
- getThread as _getThread,
31
- listThreadsByUserId as _listThreadsByUserId,
32
- updateThread as _updateThread,
33
- } from "./threads.js";
34
+ import { schema, v } from "./schema.js";
34
35
  import { insertVector, searchVectors } from "./vector/index.js";
35
36
  import {
36
- type VectorDimension,
37
- VectorDimensions,
37
+ validateVectorDimension,
38
38
  type VectorTableId,
39
39
  vVectorId,
40
40
  } from "./vector/tables.js";
41
41
  import { changeRefcount } from "./files.js";
42
+ import { getStreamingMessagesWithMetadata, finishHandler } from "./streams.js";
43
+ import { partial } from "convex-helpers/validators";
42
44
 
43
45
  function publicMessage(message: Doc<"messages">): MessageDoc {
44
46
  return omit(message, ["parentMessageId", "stepId", "files"]);
@@ -58,9 +60,7 @@ export async function deleteMessage(
58
60
  }
59
61
 
60
62
  export const deleteByIds = mutation({
61
- args: {
62
- messageIds: v.array(v.id("messages")),
63
- },
63
+ args: { messageIds: v.array(v.id("messages")) },
64
64
  returns: v.array(v.id("messages")),
65
65
  handler: async (ctx, args) => {
66
66
  const deletedMessageIds = await Promise.all(
@@ -94,13 +94,20 @@ export const deleteByOrder = mutation({
94
94
  lastOrder: v.optional(v.number()),
95
95
  lastStepOrder: v.optional(v.number()),
96
96
  }),
97
- handler: async (ctx, args) => {
98
- const messages = await orderedMessagesStream(
99
- ctx,
100
- args.threadId,
101
- "asc",
102
- args.startOrder,
103
- )
97
+ handler: async (
98
+ ctx,
99
+ args,
100
+ ): Promise<{
101
+ isDone: boolean;
102
+ lastOrder?: number;
103
+ lastStepOrder?: number;
104
+ }> => {
105
+ const messages = await orderedMessagesStream(ctx, {
106
+ threadId: args.threadId,
107
+ sortOrder: "asc",
108
+ startOrder: args.startOrder,
109
+ startOrderBound: "gte",
110
+ })
104
111
  .narrow({
105
112
  lowerBound: args.startStepOrder
106
113
  ? [args.startOrder, args.startStepOrder]
@@ -127,16 +134,21 @@ const addMessagesArgs = {
127
134
  promptMessageId: v.optional(v.id("messages")),
128
135
  agentName: v.optional(v.string()),
129
136
  messages: v.array(vMessageWithMetadataInternal),
130
- embeddings: v.optional(vMessageEmbeddings),
131
- pending: v.optional(v.boolean()),
137
+ embeddings: v.optional(vMessageEmbeddingsWithDimension),
132
138
  failPendingSteps: v.optional(v.boolean()),
139
+ // A pending message to update. If the pending message failed, abort.
140
+ pendingMessageId: v.optional(v.id("messages")),
141
+ // if set to true, these messages will not show up in text or vector search
142
+ // results for the userId
143
+ hideFromUserIdSearch: v.optional(v.boolean()),
144
+ // If provided, finish this stream atomically with the message save.
145
+ // This prevents UI flickering from separate mutations (issue #181).
146
+ finishStreamId: v.optional(v.id("streamingMessages")),
133
147
  };
134
148
  export const addMessages = mutation({
135
149
  args: addMessagesArgs,
136
150
  handler: addMessagesHandler,
137
- returns: v.object({
138
- messages: v.array(vMessageDoc),
139
- }),
151
+ returns: v.object({ messages: v.array(vMessageDoc) }),
140
152
  });
141
153
  async function addMessagesHandler(
142
154
  ctx: MutationCtx,
@@ -152,12 +164,15 @@ async function addMessagesHandler(
152
164
  const {
153
165
  embeddings,
154
166
  failPendingSteps,
155
- pending,
167
+ // Destructured separately to exclude from `...rest` (used in addMessages args, not message fields)
168
+ finishStreamId,
156
169
  messages,
157
170
  promptMessageId,
171
+ pendingMessageId,
172
+ hideFromUserIdSearch,
158
173
  ...rest
159
174
  } = args;
160
- const parentMessage = promptMessageId && (await ctx.db.get(promptMessageId));
175
+ const promptMessage = promptMessageId && (await ctx.db.get(promptMessageId));
161
176
  if (failPendingSteps) {
162
177
  assert(args.threadId, "threadId is required to fail pending steps");
163
178
  const pendingMessages = await ctx.db
@@ -165,30 +180,41 @@ async function addMessagesHandler(
165
180
  .withIndex("threadId_status_tool_order_stepOrder", (q) =>
166
181
  q.eq("threadId", threadId).eq("status", "pending"),
167
182
  )
168
- .collect();
183
+ .order("desc")
184
+ .take(100);
169
185
  await Promise.all(
170
186
  pendingMessages
171
- .filter((m) => !parentMessage || m.order === parentMessage.order)
172
- .map((m) =>
173
- ctx.db.patch(m._id, { status: "failed", error: "Restarting" }),
174
- ),
187
+ .filter((m) => !promptMessage || m.order === promptMessage.order)
188
+ .filter((m) => !pendingMessageId || m._id !== pendingMessageId)
189
+ .map(async (m) => {
190
+ if (m.embeddingId) {
191
+ await ctx.db.delete(m.embeddingId);
192
+ }
193
+ await ctx.db.patch(m._id, {
194
+ status: "failed",
195
+ error: "Restarting",
196
+ embeddingId: undefined,
197
+ });
198
+ }),
175
199
  );
176
200
  }
177
201
  let order, stepOrder;
178
202
  let fail = false;
203
+ let error: string | undefined;
179
204
  if (promptMessageId) {
180
- assert(parentMessage, `Parent message ${promptMessageId} not found`);
181
- if (parentMessage.status === "failed") {
205
+ assert(promptMessage, `Parent message ${promptMessageId} not found`);
206
+ if (promptMessage.status === "failed") {
182
207
  fail = true;
208
+ error = promptMessage.error ?? error ?? "The prompt message failed";
183
209
  }
184
- order = parentMessage.order;
210
+ order = promptMessage.order;
185
211
  // Defend against there being existing messages with this parent.
186
212
  const maxMessage = await getMaxMessage(ctx, threadId, order);
187
- stepOrder = maxMessage?.stepOrder ?? parentMessage.stepOrder;
213
+ stepOrder = maxMessage?.stepOrder ?? promptMessage.stepOrder;
188
214
  } else {
189
215
  const maxMessage = await getMaxMessage(ctx, threadId);
190
- order = maxMessage ? maxMessage.order + 1 : 0;
191
- stepOrder = -1;
216
+ order = maxMessage?.order ?? -1;
217
+ stepOrder = maxMessage?.stepOrder ?? -1;
192
218
  }
193
219
  const toReturn: Doc<"messages">[] = [];
194
220
  if (embeddings) {
@@ -200,41 +226,93 @@ async function addMessagesHandler(
200
226
  for (let i = 0; i < messages.length; i++) {
201
227
  const message = messages[i];
202
228
  let embeddingId: VectorTableId | undefined;
203
- if (embeddings && embeddings.vectors[i]) {
229
+ if (
230
+ embeddings &&
231
+ embeddings.vectors[i] &&
232
+ !fail &&
233
+ message.status !== "failed"
234
+ ) {
204
235
  embeddingId = await insertVector(ctx, embeddings.dimension, {
205
236
  vector: embeddings.vectors[i]!,
206
237
  model: embeddings.model,
207
238
  table: "messages",
208
- userId,
239
+ userId: hideFromUserIdSearch ? undefined : userId,
209
240
  threadId,
210
241
  });
211
242
  }
212
- stepOrder++;
213
- const messageId = await ctx.db.insert("messages", {
243
+ const messageDoc = {
214
244
  ...rest,
215
245
  ...message,
216
246
  embeddingId,
217
247
  parentMessageId: promptMessageId,
218
248
  userId,
219
- order,
220
249
  tool: isTool(message.message),
221
- text: extractText(message.message),
222
- status: fail ? "failed" : pending ? "pending" : "success",
223
- error: fail ? "Parent message failed" : undefined,
250
+ text: hideFromUserIdSearch ? undefined : extractText(message.message),
251
+ status: fail ? "failed" : (message.status ?? "success"),
252
+ error: fail ? error : message.error,
253
+ } satisfies Omit<
254
+ WithoutSystemFields<Doc<"messages">>,
255
+ "order" | "stepOrder"
256
+ >;
257
+ // If there is a pending message, we replace that one with the first message
258
+ // and subsequent ones will follow the regular order/subOrder advancement.
259
+ if (i === 0 && pendingMessageId) {
260
+ const pendingMessage = await ctx.db.get(pendingMessageId);
261
+ assert(pendingMessage, `Pending msg ${pendingMessageId} not found`);
262
+ if (pendingMessage.status === "failed") {
263
+ fail = true;
264
+ error =
265
+ `Trying to update a message that failed: ${pendingMessageId}, ` +
266
+ `error: ${pendingMessage.error ?? error}`;
267
+ messageDoc.status = "failed";
268
+ messageDoc.error = error;
269
+ }
270
+ if (message.fileIds) {
271
+ await changeRefcount(
272
+ ctx,
273
+ pendingMessage.fileIds ?? [],
274
+ message.fileIds,
275
+ );
276
+ }
277
+ await ctx.db.replace(pendingMessage._id, {
278
+ ...messageDoc,
279
+ order: pendingMessage.order,
280
+ stepOrder: pendingMessage.stepOrder,
281
+ });
282
+ toReturn.push((await ctx.db.get(pendingMessage._id))!);
283
+ continue;
284
+ }
285
+ if (message.message.role === "user") {
286
+ if (promptMessage && promptMessage.order === order) {
287
+ // see if there's a later message than the parent message order
288
+ const maxMessage = await getMaxMessage(ctx, threadId);
289
+ order = (maxMessage?.order ?? order) + 1;
290
+ } else {
291
+ order++;
292
+ }
293
+ stepOrder = 0;
294
+ } else {
295
+ if (order < 0) {
296
+ order = 0;
297
+ }
298
+ stepOrder++;
299
+ }
300
+ const messageId = await ctx.db.insert("messages", {
301
+ ...messageDoc,
302
+ order,
224
303
  stepOrder,
225
304
  });
226
- // Let's just not set the id field and have it set only in explicit cases.
227
- // if (!message.id) {
228
- // await ctx.db.patch(messageId, {
229
- // id: messageId,
230
- // });
231
- // }
232
305
  if (message.fileIds) {
233
306
  await changeRefcount(ctx, [], message.fileIds);
234
307
  }
235
308
  // TODO: delete the associated stream data for the order/stepOrder
236
309
  toReturn.push((await ctx.db.get(messageId))!);
237
310
  }
311
+ // Atomically finish the stream if requested, preventing UI flickering
312
+ // from separate mutations for message save and stream finish (issue #181).
313
+ if (finishStreamId) {
314
+ await finishHandler(ctx, { streamId: finishStreamId });
315
+ }
238
316
  return { messages: toReturn.map(publicMessage) };
239
317
  }
240
318
 
@@ -244,14 +322,22 @@ export async function getMaxMessage(
244
322
  threadId: Id<"threads">,
245
323
  order?: number,
246
324
  ) {
247
- return orderedMessagesStream(ctx, threadId, "desc", order).first();
325
+ return orderedMessagesStream(ctx, {
326
+ threadId,
327
+ sortOrder: "desc",
328
+ startOrder: order,
329
+ startOrderBound: "eq",
330
+ }).first();
248
331
  }
249
332
 
250
333
  function orderedMessagesStream(
251
334
  ctx: QueryCtx,
252
- threadId: Id<"threads">,
253
- sortOrder: "asc" | "desc",
254
- order?: number,
335
+ args: {
336
+ threadId: Id<"threads">;
337
+ sortOrder: "asc" | "desc";
338
+ startOrder?: number;
339
+ startOrderBound?: "gte" | "eq";
340
+ },
255
341
  ) {
256
342
  return mergedStream(
257
343
  [true, false].flatMap((tool) =>
@@ -260,66 +346,96 @@ function orderedMessagesStream(
260
346
  .query("messages")
261
347
  .withIndex("threadId_status_tool_order_stepOrder", (q) => {
262
348
  const qq = q
263
- .eq("threadId", threadId)
349
+ .eq("threadId", args.threadId)
264
350
  .eq("status", status)
265
351
  .eq("tool", tool);
266
- if (order) {
267
- return qq.eq("order", order);
352
+ if (args.startOrder !== undefined) {
353
+ if (args.startOrderBound === "gte") {
354
+ return qq.gte("order", args.startOrder);
355
+ } else {
356
+ return qq.eq("order", args.startOrder);
357
+ }
268
358
  }
269
359
  return qq;
270
360
  })
271
- .order(sortOrder),
361
+ .order(args.sortOrder),
272
362
  ),
273
363
  ),
274
364
  ["order", "stepOrder"],
275
365
  );
276
366
  }
277
367
 
278
- export const rollbackMessage = mutation({
368
+ export const finalizeMessage = mutation({
279
369
  args: {
280
370
  messageId: v.id("messages"),
281
- error: v.optional(v.string()),
371
+ result: v.union(
372
+ v.object({ status: v.literal("success") }),
373
+ v.object({ status: v.literal("failed"), error: v.string() }),
374
+ ),
282
375
  },
283
376
  returns: v.null(),
284
- handler: async (ctx, { messageId, error }) => {
377
+ handler: async (ctx, { messageId, result }) => {
285
378
  const message = await ctx.db.get(messageId);
286
379
  assert(message, `Message ${messageId} not found`);
287
- const messages = await orderedMessagesStream(
288
- ctx,
289
- message.threadId,
290
- "asc",
291
- message.order,
292
- ).collect();
293
- for (const m of messages) {
294
- if (m.status === "pending") {
295
- await ctx.db.patch(m._id, { status: "failed", error });
380
+ if (message.status !== "pending") {
381
+ console.debug(
382
+ "Trying to finalize a message that's already",
383
+ message.status,
384
+ );
385
+ return;
386
+ }
387
+ // See if we can add any in-progress data
388
+ if (!message.message?.content.length) {
389
+ const messages = await getStreamingMessagesWithMetadata(
390
+ ctx,
391
+ message,
392
+ result,
393
+ );
394
+ if (messages.length > 0) {
395
+ await addMessagesHandler(ctx, {
396
+ messages,
397
+ threadId: message.threadId,
398
+ agentName: message.agentName,
399
+ failPendingSteps: false,
400
+ pendingMessageId: messageId,
401
+ userId: message.userId,
402
+ embeddings: undefined,
403
+ });
404
+ return;
296
405
  }
297
406
  }
298
-
299
- await ctx.db.patch(messageId, {
300
- status: "failed",
301
- error: error,
302
- });
303
- },
304
- });
305
-
306
- export const commitMessage = mutation({
307
- args: {
308
- messageId: v.id("messages"),
407
+ if (result.status === "failed") {
408
+ if (message.embeddingId) {
409
+ await ctx.db.delete(message.embeddingId);
410
+ }
411
+ await ctx.db.patch(messageId, {
412
+ status: "failed",
413
+ error: result.error,
414
+ embeddingId: undefined,
415
+ });
416
+ } else {
417
+ await ctx.db.patch(messageId, { status: "success" });
418
+ }
309
419
  },
310
- returns: v.null(),
311
- handler: commitMessageHandler,
312
420
  });
313
421
 
314
422
  export const updateMessage = mutation({
315
423
  args: {
316
424
  messageId: v.id("messages"),
317
- patch: v.object({
318
- message: v.optional(vMessageDoc.fields.message),
319
- fileIds: v.optional(v.array(v.id("files"))),
320
- status: v.optional(vMessageStatus),
321
- error: v.optional(v.string()),
322
- }),
425
+ patch: v.object(
426
+ partial(
427
+ pick(schema.tables.messages.validator.fields, [
428
+ "message",
429
+ "fileIds",
430
+ "status",
431
+ "error",
432
+ "model",
433
+ "provider",
434
+ "providerOptions",
435
+ "finishReason",
436
+ ]),
437
+ ),
438
+ ),
323
439
  },
324
440
  returns: vMessageDoc,
325
441
  handler: async (ctx, args) => {
@@ -330,9 +446,7 @@ export const updateMessage = mutation({
330
446
  await changeRefcount(ctx, message.fileIds ?? [], args.patch.fileIds);
331
447
  }
332
448
 
333
- const patch: Partial<Doc<"messages">> = {
334
- ...args.patch,
335
- };
449
+ const patch: Partial<Doc<"messages">> = { ...args.patch };
336
450
 
337
451
  if (args.patch.message !== undefined) {
338
452
  patch.message = args.patch.message;
@@ -340,102 +454,234 @@ export const updateMessage = mutation({
340
454
  patch.text = extractText(args.patch.message);
341
455
  }
342
456
 
457
+ if (args.patch.status === "failed") {
458
+ if (message.embeddingId) {
459
+ await ctx.db.delete(message.embeddingId);
460
+ }
461
+ patch.embeddingId = undefined;
462
+ }
463
+
343
464
  await ctx.db.patch(args.messageId, patch);
344
465
  return publicMessage((await ctx.db.get(args.messageId))!);
345
466
  },
346
467
  });
347
468
 
348
- async function commitMessageHandler(
349
- ctx: MutationCtx,
350
- { messageId }: { messageId: Id<"messages"> },
351
- ) {
352
- const message = await ctx.db.get(messageId);
353
- assert(message, `Message ${messageId} not found`);
469
+ const cloneMessageArgs = {
470
+ sourceThreadId: v.id("threads"),
471
+ targetThreadId: v.id("threads"),
472
+ // defaults to false, so searching for a message by userId will not find
473
+ // these copies
474
+ copyUserIdForVectorSearch: v.optional(v.boolean()),
475
+ // defaults to false, so tool calls & responses will be copied
476
+ excludeToolMessages: v.optional(v.boolean()),
477
+ // defaults to copying all messages, but you could just copy success messages.
478
+ statuses: v.optional(v.array(vMessageStatus)),
479
+ // stop at this message id
480
+ upToAndIncludingMessageId: v.optional(v.id("messages")),
481
+ // defaults to 0. the messages will be inserted starting at this order.
482
+ insertAtOrder: v.optional(v.number()),
483
+ };
484
+ export const cloneMessageBatch = internalMutation({
485
+ args: {
486
+ ...cloneMessageArgs,
487
+ paginationOpts: paginationOptsValidator,
488
+ },
489
+ handler: async (
490
+ ctx,
491
+ args,
492
+ ): Promise<{
493
+ numCopied: number;
494
+ continueCursor: string;
495
+ isDone: boolean;
496
+ }> => {
497
+ const orderOffset = args.insertAtOrder ?? 0;
498
+ const result = await listMessagesByThreadIdHandler(ctx, {
499
+ threadId: args.sourceThreadId,
500
+ excludeToolMessages: args.excludeToolMessages,
501
+ order: "desc",
502
+ paginationOpts: args.paginationOpts,
503
+ statuses: args.statuses,
504
+ upToAndIncludingMessageId: args.upToAndIncludingMessageId,
505
+ });
354
506
 
355
- const order = message.order!;
356
- const messages = await mergedStream(
357
- [true, false].map((tool) =>
358
- stream(ctx.db, schema)
359
- .query("messages")
360
- .withIndex("threadId_status_tool_order_stepOrder", (q) =>
361
- q
362
- .eq("threadId", message.threadId)
363
- .eq("status", "pending")
364
- .eq("tool", tool)
365
- .eq("order", order),
366
- ),
367
- ),
368
- ["order", "stepOrder"],
369
- ).collect();
370
- for (const message of messages) {
371
- await ctx.db.patch(message._id, { status: "success" });
372
- }
373
- }
507
+ const existing =
508
+ result.page.length === 0
509
+ ? []
510
+ : await mergedStream(
511
+ [true, false].flatMap((tool) =>
512
+ messageStatuses.map((status) =>
513
+ stream(ctx.db, schema)
514
+ .query("messages")
515
+ .withIndex("threadId_status_tool_order_stepOrder", (q) =>
516
+ q
517
+ .eq("threadId", args.targetThreadId)
518
+ .eq("status", status)
519
+ .eq("tool", tool)
520
+ .gte("order", result.page[0].order)
521
+ .lte("order", result.page[result.page.length - 1].order),
522
+ ),
523
+ ),
524
+ ),
525
+ ["order", "stepOrder"],
526
+ ).collect();
374
527
 
375
- export const listMessagesByThreadId = query({
528
+ await Promise.all(
529
+ result.page
530
+ .filter(
531
+ (m) =>
532
+ !existing.some(
533
+ (e) => e.order === m.order && e.stepOrder === m.stepOrder,
534
+ ),
535
+ )
536
+ .map(async (m) => {
537
+ // update file refs
538
+ if (m.fileIds) {
539
+ await changeRefcount(ctx, [], m.fileIds);
540
+ }
541
+ let embeddingId: VectorTableId | undefined = undefined;
542
+ if (m.embeddingId) {
543
+ const vector = await ctx.db.get(m.embeddingId);
544
+ assert(vector, `Vector ${m.embeddingId} not found`);
545
+ const dimension = vector.vector.length;
546
+ validateVectorDimension(dimension);
547
+ embeddingId = await insertVector(ctx, dimension, {
548
+ ...pick(vector, ["model", "table", "vector"]),
549
+ userId: args.copyUserIdForVectorSearch
550
+ ? vector.userId
551
+ : undefined,
552
+ threadId: args.targetThreadId,
553
+ });
554
+ }
555
+ await ctx.db.insert("messages", {
556
+ ...omit(m, [
557
+ "_id",
558
+ "_creationTime",
559
+ "threadId",
560
+ "order",
561
+ "embeddingId",
562
+ ]),
563
+ embeddingId,
564
+ threadId: args.targetThreadId,
565
+ order: orderOffset + m.order,
566
+ });
567
+ }),
568
+ );
569
+ return {
570
+ numCopied: result.page.length,
571
+ continueCursor: result.continueCursor,
572
+ isDone: result.isDone,
573
+ };
574
+ },
575
+ });
576
+
577
+ export const cloneThread = action({
376
578
  args: {
377
- threadId: v.id("threads"),
378
- excludeToolMessages: v.optional(v.boolean()),
379
- /** What order to sort the messages in. To get the latest, use "desc". */
380
- order: v.union(v.literal("asc"), v.literal("desc")),
381
- paginationOpts: v.optional(paginationOptsValidator),
382
- statuses: v.optional(v.array(vMessageStatus)),
383
- upToAndIncludingMessageId: v.optional(v.id("messages")),
579
+ ...cloneMessageArgs,
580
+ batchSize: v.optional(v.number()),
581
+ // how many messages to copy
582
+ limit: v.optional(v.number()),
384
583
  },
584
+ returns: v.number(),
385
585
  handler: async (ctx, args) => {
386
- const statuses =
387
- args.statuses ?? vMessageStatus.members.map((m) => m.value);
388
- const last =
389
- args.upToAndIncludingMessageId &&
390
- (await ctx.db.get(args.upToAndIncludingMessageId));
391
- assert(
392
- !last || last.threadId === args.threadId,
393
- "upToAndIncludingMessageId must be a message in the thread",
394
- );
395
- const toolOptions = args.excludeToolMessages ? [false] : [true, false];
396
- const order = args.order ?? "desc";
397
- const streams = toolOptions.flatMap((tool) =>
398
- statuses.map((status) =>
399
- stream(ctx.db, schema)
400
- .query("messages")
401
- .withIndex("threadId_status_tool_order_stepOrder", (q) => {
402
- const qq = q
403
- .eq("threadId", args.threadId)
404
- .eq("status", status)
405
- .eq("tool", tool);
406
- if (last) {
407
- return qq.lte("order", last.order);
408
- }
409
- return qq;
410
- })
411
- .order(order)
412
- .filterWith(
413
- // We allow all messages on the same order.
414
- async (m) =>
415
- !last || m.order < last.order || m.order === last.order,
416
- ),
417
- ),
418
- );
419
- const messages = await mergedStream(streams, [
420
- "order",
421
- "stepOrder",
422
- ]).paginate(
423
- args.paginationOpts ?? {
424
- numItems: DEFAULT_RECENT_MESSAGES,
425
- cursor: null,
426
- },
427
- );
586
+ let cursor: string | null = null;
587
+ let copiedSoFar = 0;
588
+ while (copiedSoFar < (args.limit ?? Infinity)) {
589
+ const numToCopy = Math.min(
590
+ args.batchSize ?? DEFAULT_RECENT_MESSAGES,
591
+ args.limit ?? Infinity - copiedSoFar,
592
+ );
593
+ const result: {
594
+ numCopied: number;
595
+ continueCursor: string;
596
+ isDone: boolean;
597
+ } = await ctx.runMutation(internal.messages.cloneMessageBatch, {
598
+ ...args,
599
+ paginationOpts: {
600
+ cursor,
601
+ numItems: numToCopy,
602
+ },
603
+ });
604
+ copiedSoFar += result.numCopied;
605
+ cursor = result.continueCursor;
606
+ if (result.isDone) {
607
+ break;
608
+ }
609
+ }
610
+ return copiedSoFar;
611
+ },
612
+ });
613
+
614
+ export const listMessagesByThreadIdArgs = {
615
+ threadId: v.id("threads"),
616
+ excludeToolMessages: v.optional(v.boolean()),
617
+ /** What order to sort the messages in. To get the latest, use "desc". */
618
+ order: v.union(v.literal("asc"), v.literal("desc")),
619
+ paginationOpts: v.optional(paginationOptsValidator),
620
+ statuses: v.optional(v.array(vMessageStatus)),
621
+ upToAndIncludingMessageId: v.optional(v.id("messages")),
622
+ };
623
+ export const listMessagesByThreadId = query({
624
+ args: listMessagesByThreadIdArgs,
625
+ handler: async (ctx, args) => {
626
+ const messages = await listMessagesByThreadIdHandler(ctx, args);
428
627
  return { ...messages, page: messages.page.map(publicMessage) };
429
628
  },
430
629
  returns: vPaginationResult(vMessageDoc),
431
630
  });
432
631
 
632
+ async function listMessagesByThreadIdHandler(
633
+ ctx: QueryCtx,
634
+ args: ObjectType<typeof listMessagesByThreadIdArgs>,
635
+ ) {
636
+ const statuses = args.statuses ?? vMessageStatus.members.map((m) => m.value);
637
+ const last =
638
+ args.upToAndIncludingMessageId &&
639
+ (await ctx.db.get(args.upToAndIncludingMessageId));
640
+ assert(
641
+ !last || last.threadId === args.threadId,
642
+ "upToAndIncludingMessageId must be a message in the thread",
643
+ );
644
+ const toolOptions = args.excludeToolMessages ? [false] : [true, false];
645
+ const order = args.order ?? "desc";
646
+ const streams = toolOptions.flatMap((tool) =>
647
+ statuses.map((status) =>
648
+ stream(ctx.db, schema)
649
+ .query("messages")
650
+ .withIndex("threadId_status_tool_order_stepOrder", (q) => {
651
+ const qq = q
652
+ .eq("threadId", args.threadId)
653
+ .eq("status", status)
654
+ .eq("tool", tool);
655
+ if (last) {
656
+ return qq.lte("order", last.order);
657
+ }
658
+ return qq;
659
+ })
660
+ .order(order)
661
+ .filterWith(
662
+ // We allow all messages on the same order.
663
+ async (m) => !last || m.order <= last.order,
664
+ ),
665
+ ),
666
+ );
667
+ const messages = await mergedStream(streams, ["order", "stepOrder"]).paginate(
668
+ args.paginationOpts ?? {
669
+ numItems: DEFAULT_RECENT_MESSAGES,
670
+ cursor: null,
671
+ },
672
+ );
673
+ if (messages.page.length === 0) {
674
+ messages.isDone = true;
675
+ }
676
+ return messages;
677
+ }
678
+
433
679
  export const getMessagesByIds = query({
434
- args: {
435
- messageIds: v.array(v.id("messages")),
436
- },
680
+ args: { messageIds: v.array(v.id("messages")) },
437
681
  handler: async (ctx, args) => {
438
- return await Promise.all(args.messageIds.map((id) => ctx.db.get(id)));
682
+ return (await Promise.all(args.messageIds.map((id) => ctx.db.get(id)))).map(
683
+ (m) => (m ? publicMessage(m) : null),
684
+ );
439
685
  },
440
686
  returns: v.array(v.union(v.null(), vMessageDoc)),
441
687
  });
@@ -444,10 +690,12 @@ export const searchMessages = action({
444
690
  args: {
445
691
  threadId: v.optional(v.id("threads")),
446
692
  searchAllMessagesForUserId: v.optional(v.string()),
447
- beforeMessageId: v.optional(v.id("messages")),
693
+ targetMessageId: v.optional(v.id("messages")),
448
694
  embedding: v.optional(v.array(v.number())),
449
695
  embeddingModel: v.optional(v.string()),
450
696
  text: v.optional(v.string()),
697
+ textSearch: v.optional(v.boolean()),
698
+ vectorSearch: v.optional(v.boolean()),
451
699
  limit: v.number(),
452
700
  vectorScoreThreshold: v.optional(v.number()),
453
701
  messageRange: v.optional(
@@ -462,24 +710,38 @@ export const searchMessages = action({
462
710
  );
463
711
  const limit = args.limit;
464
712
  let textSearchMessages: MessageDoc[] | undefined;
465
- if (args.text) {
713
+ if (args.textSearch) {
466
714
  textSearchMessages = await ctx.runQuery(api.messages.textSearch, {
467
715
  searchAllMessagesForUserId: args.searchAllMessagesForUserId,
468
716
  threadId: args.threadId,
717
+ targetMessageId: args.targetMessageId,
469
718
  text: args.text,
470
719
  limit,
471
- beforeMessageId: args.beforeMessageId,
472
720
  });
473
721
  }
474
- if (args.embedding) {
475
- const dimension = args.embedding.length as VectorDimension;
476
- if (!VectorDimensions.includes(dimension)) {
477
- throw new Error(`Unsupported embedding dimension: ${dimension}`);
722
+ if (args.vectorSearch) {
723
+ let embedding = args.embedding;
724
+ let model = args.embeddingModel;
725
+ if (!embedding) {
726
+ if (args.targetMessageId) {
727
+ const target = await ctx.runQuery(
728
+ api.messages.getMessageSearchFields,
729
+ {
730
+ messageId: args.targetMessageId,
731
+ },
732
+ );
733
+ assert(target, "Target message embedding not found.");
734
+ embedding = target.embedding;
735
+ model = target.embeddingModel;
736
+ }
478
737
  }
738
+ assert(embedding && model, "Embedding missing");
739
+ const dimension = embedding.length;
740
+ validateVectorDimension(dimension);
479
741
  const vectors = (
480
- await searchVectors(ctx, args.embedding, {
742
+ await searchVectors(ctx, embedding, {
481
743
  dimension,
482
- model: args.embeddingModel ?? "unknown",
744
+ model,
483
745
  table: "messages",
484
746
  searchAllMessagesForUserId: args.searchAllMessagesForUserId,
485
747
  threadId: args.threadId,
@@ -508,7 +770,7 @@ export const searchMessages = action({
508
770
  (m) => !embeddingIds.includes(m.embeddingId! as VectorTableId),
509
771
  ),
510
772
  messageRange: args.messageRange ?? DEFAULT_MESSAGE_RANGE,
511
- beforeMessageId: args.beforeMessageId,
773
+ beforeMessageId: args.targetMessageId,
512
774
  limit,
513
775
  },
514
776
  );
@@ -572,11 +834,11 @@ export const _fetchSearchMessages = internalQuery({
572
834
  .map(publicMessage);
573
835
  messages.push(...(args.textSearchMessages ?? []));
574
836
  // TODO: prioritize more recent messages
575
- messages.sort((a, b) => a.order! - b.order!);
837
+ messages = sorted(messages);
576
838
  messages = messages.slice(0, args.limit);
577
839
  // Fetch the surrounding messages
578
840
  if (!threadId) {
579
- return messages.sort((a, b) => a.order - b.order);
841
+ return messages;
580
842
  }
581
843
  const included: Record<string, Set<number>> = {};
582
844
  for (const m of messages) {
@@ -629,7 +891,7 @@ export const _fetchSearchMessages = internalQuery({
629
891
  messages.push(publicMessage(r));
630
892
  }
631
893
  }
632
- return messages.sort((a, b) => a.order - b.order);
894
+ return sorted(messages);
633
895
  },
634
896
  });
635
897
 
@@ -639,26 +901,29 @@ export const textSearch = query({
639
901
  args: {
640
902
  threadId: v.optional(v.id("threads")),
641
903
  searchAllMessagesForUserId: v.optional(v.string()),
642
- text: v.string(),
904
+ text: v.optional(v.string()),
905
+ targetMessageId: v.optional(v.id("messages")),
643
906
  limit: v.number(),
644
- beforeMessageId: v.optional(v.id("messages")),
645
907
  },
646
908
  handler: async (ctx, args) => {
647
909
  assert(
648
910
  args.searchAllMessagesForUserId || args.threadId,
649
911
  "Specify userId or threadId",
650
912
  );
651
- const beforeMessage =
652
- args.beforeMessageId && (await ctx.db.get(args.beforeMessageId));
653
- const order = beforeMessage?.order;
913
+ const targetMessage =
914
+ args.targetMessageId && (await ctx.db.get(args.targetMessageId));
915
+ const order = targetMessage?.order;
916
+ const text = args.text || targetMessage?.text;
917
+ if (!text) {
918
+ console.warn("No text to search", targetMessage, args.text);
919
+ return [];
920
+ }
654
921
  const messages = await ctx.db
655
922
  .query("messages")
656
923
  .withSearchIndex("text_search", (q) =>
657
924
  args.searchAllMessagesForUserId
658
- ? q
659
- .search("text", args.text)
660
- .eq("userId", args.searchAllMessagesForUserId)
661
- : q.search("text", args.text).eq("threadId", args.threadId!),
925
+ ? q.search("text", text).eq("userId", args.searchAllMessagesForUserId)
926
+ : q.search("text", text).eq("threadId", args.threadId!),
662
927
  )
663
928
  // Just in case tool messages slip through
664
929
  .filter((q) => {
@@ -672,12 +937,46 @@ export const textSearch = query({
672
937
  return messages
673
938
  .filter(
674
939
  (m) =>
675
- !beforeMessage ||
676
- m.order < beforeMessage.order ||
677
- (m.order === beforeMessage.order &&
678
- m.stepOrder < beforeMessage.stepOrder),
940
+ !targetMessage ||
941
+ m.order < targetMessage.order ||
942
+ (m.order === targetMessage.order &&
943
+ m.stepOrder < targetMessage.stepOrder),
679
944
  )
680
945
  .map(publicMessage);
681
946
  },
682
947
  returns: v.array(vMessageDoc),
683
948
  });
949
+
950
+ export const getMessageSearchFields = query({
951
+ args: {
952
+ messageId: v.id("messages"),
953
+ },
954
+ returns: v.object({
955
+ text: v.optional(v.string()),
956
+ embedding: v.optional(v.array(v.number())),
957
+ embeddingModel: v.optional(v.string()),
958
+ }),
959
+ handler: async (
960
+ ctx,
961
+ args,
962
+ ): Promise<{
963
+ text?: string | undefined;
964
+ embedding?: number[] | undefined;
965
+ embeddingModel?: string | undefined;
966
+ }> => {
967
+ const message = await ctx.db.get(args.messageId);
968
+ const text = message?.text;
969
+ let embedding = undefined;
970
+ let embeddingModel = undefined;
971
+ if (message?.embeddingId) {
972
+ const target = await ctx.db.get(message.embeddingId);
973
+ embedding = target?.vector;
974
+ embeddingModel = target?.model;
975
+ }
976
+ return {
977
+ text,
978
+ embedding,
979
+ embeddingModel,
980
+ };
981
+ },
982
+ });