@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,25 +1,55 @@
1
- import type {
2
- AgentComponent,
3
- ContextOptions,
4
- RunActionCtx,
5
- RunQueryCtx,
6
- } from "./types.js";
7
- import type { MessageDoc } from "../component/schema.js";
8
- import type { ModelMessage } from "ai";
1
+ import {
2
+ embedMany as embedMany_,
3
+ type EmbeddingModel,
4
+ type ModelMessage,
5
+ } from "ai";
9
6
  import { assert } from "convex-helpers";
7
+ import type { MessageDoc } from "../validators.js";
8
+ import {
9
+ validateVectorDimension,
10
+ type VectorDimension,
11
+ } from "../component/vector/tables.js";
10
12
  import {
11
13
  DEFAULT_MESSAGE_RANGE,
12
14
  DEFAULT_RECENT_MESSAGES,
13
15
  extractText,
16
+ getModelName,
17
+ getProviderName,
18
+ isTool,
19
+ sorted,
14
20
  } from "../shared.js";
15
21
  import type { Message } from "../validators.js";
22
+ import type {
23
+ AgentComponent,
24
+ Config,
25
+ ContextOptions,
26
+ Options,
27
+ QueryCtx,
28
+ MutationCtx,
29
+ ActionCtx,
30
+ } from "./types.js";
31
+ import { inlineMessagesFiles } from "./files.js";
32
+ import { docsToModelMessages, toModelMessage } from "../mapping.js";
16
33
 
17
34
  const DEFAULT_VECTOR_SCORE_THRESHOLD = 0.0;
35
+ // 10k characters should be more than enough for most cases, and stays under
36
+ // the 8k token limit for some models.
37
+ const MAX_EMBEDDING_TEXT_LENGTH = 10_000;
18
38
 
19
- export type GetEmbedding = (text: string) => Promise<{
20
- embedding: number[];
21
- embeddingModel: string;
22
- }>;
39
+ export type GetEmbedding = (text: string) => Promise<
40
+ | {
41
+ embedding: number[];
42
+ /** @deprecated Use embeddingModel instead. */
43
+ textEmbeddingModel: string | EmbeddingModel;
44
+ embeddingModel?: string | EmbeddingModel;
45
+ }
46
+ | {
47
+ embedding: number[];
48
+ /** @deprecated Use embeddingModel instead. */
49
+ textEmbeddingModel?: string | EmbeddingModel;
50
+ embeddingModel: string | EmbeddingModel;
51
+ }
52
+ >;
23
53
 
24
54
  /**
25
55
  * Fetch the context messages for a thread.
@@ -30,31 +60,78 @@ export type GetEmbedding = (text: string) => Promise<{
30
60
  * @returns
31
61
  */
32
62
  export async function fetchContextMessages(
33
- ctx: RunQueryCtx | RunActionCtx,
63
+ ctx: QueryCtx | MutationCtx | ActionCtx,
34
64
  component: AgentComponent,
35
65
  args: {
36
66
  userId: string | undefined;
37
67
  threadId: string | undefined;
38
- messages: (ModelMessage | Message)[];
39
68
  /**
40
- * If provided, it will search for messages up to and including this message.
41
- * Note: if this is far in the past, text and vector search results may be more
42
- * limited, as it's post-filtering the results.
69
+ * If targetMessageId is not provided, this text will be used
70
+ * for text and vector search
71
+ */
72
+ searchText?: string;
73
+ /**
74
+ * If provided, it will use this message for text/vector search (if enabled)
75
+ * and will only fetch messages up to (and including) this message's "order"
76
+ */
77
+ targetMessageId?: string;
78
+ /**
79
+ * @deprecated use searchText and targetMessageId instead
80
+ */
81
+ messages?: (ModelMessage | Message)[];
82
+ /**
83
+ * @deprecated use targetMessageId instead
43
84
  */
44
85
  upToAndIncludingMessageId?: string;
45
86
  contextOptions: ContextOptions;
46
87
  getEmbedding?: GetEmbedding;
47
88
  },
48
89
  ): Promise<MessageDoc[]> {
90
+ const { recentMessages, searchMessages } = await fetchRecentAndSearchMessages(
91
+ ctx,
92
+ component,
93
+ args,
94
+ );
95
+ return [...searchMessages, ...recentMessages];
96
+ }
97
+
98
+ export async function fetchRecentAndSearchMessages(
99
+ ctx: QueryCtx | MutationCtx | ActionCtx,
100
+ component: AgentComponent,
101
+ args: {
102
+ userId: string | undefined;
103
+ threadId: string | undefined;
104
+ /**
105
+ * If targetMessageId is not provided, this text will be used
106
+ * for text and vector search
107
+ */
108
+ searchText?: string;
109
+ /**
110
+ * If provided, it will use this message for text/vector search (if enabled)
111
+ * and will only fetch messages up to (and including) this message's "order"
112
+ */
113
+ targetMessageId?: string;
114
+ /**
115
+ * @deprecated use searchText and targetMessageId instead
116
+ */
117
+ messages?: (ModelMessage | Message)[];
118
+ /**
119
+ * @deprecated use targetMessageId instead
120
+ */
121
+ upToAndIncludingMessageId?: string;
122
+ contextOptions: ContextOptions;
123
+ getEmbedding?: GetEmbedding;
124
+ },
125
+ ): Promise<{ recentMessages: MessageDoc[]; searchMessages: MessageDoc[] }> {
49
126
  assert(args.userId || args.threadId, "Specify userId or threadId");
50
127
  const opts = args.contextOptions;
51
128
  // Fetch the latest messages from the thread
52
129
  let included: Set<string> | undefined;
53
- const contextMessages: MessageDoc[] = [];
54
- if (
55
- args.threadId &&
56
- (opts.recentMessages !== 0 || args.upToAndIncludingMessageId)
57
- ) {
130
+ let recentMessages: MessageDoc[] = [];
131
+ let searchMessages: MessageDoc[] = [];
132
+ const targetMessageId =
133
+ args.targetMessageId ?? args.upToAndIncludingMessageId;
134
+ if (args.threadId && opts.recentMessages !== 0) {
58
135
  const { page } = await ctx.runQuery(
59
136
  component.messages.listMessagesByThreadId,
60
137
  {
@@ -64,106 +141,187 @@ export async function fetchContextMessages(
64
141
  numItems: opts.recentMessages ?? DEFAULT_RECENT_MESSAGES,
65
142
  cursor: null,
66
143
  },
67
- upToAndIncludingMessageId: args.upToAndIncludingMessageId,
144
+ upToAndIncludingMessageId: targetMessageId,
68
145
  order: "desc",
69
146
  statuses: ["success"],
70
147
  },
71
148
  );
72
149
  included = new Set(page.map((m) => m._id));
73
- contextMessages.push(
74
- // Reverse since we fetched in descending order
75
- ...page.reverse(),
76
- );
150
+ recentMessages = filterOutOrphanedToolMessages(sorted(page));
77
151
  }
78
- if (opts.searchOptions?.textSearch || opts.searchOptions?.vectorSearch) {
79
- const targetMessage = contextMessages.find(
80
- (m) => m._id === args.upToAndIncludingMessageId,
81
- )?.message;
82
- const messagesToSearch = targetMessage ? [targetMessage] : args.messages;
152
+ if (
153
+ (opts.searchOptions?.textSearch || opts.searchOptions?.vectorSearch) &&
154
+ opts.searchOptions?.limit
155
+ ) {
83
156
  if (!("runAction" in ctx)) {
84
157
  throw new Error("searchUserMessages only works in an action");
85
158
  }
86
- const lastMessage = messagesToSearch.at(-1)!;
87
- assert(lastMessage, "No messages to search");
88
- const text = extractText(lastMessage);
89
- assert(text, `No text to search in message ${JSON.stringify(lastMessage)}`);
90
- assert(
91
- !args.contextOptions?.searchOptions?.vectorSearch || "runAction" in ctx,
92
- "You must do vector search from an action",
93
- );
94
- if (opts.searchOptions?.vectorSearch && !args.getEmbedding) {
95
- throw new Error(
96
- "You must provide an embedding and embeddingModel to use vector search",
97
- );
159
+ let text = args.searchText;
160
+ let embedding: number[] | undefined;
161
+ let embeddingModel: string | undefined;
162
+ if (!text) {
163
+ if (targetMessageId) {
164
+ const targetMessage = recentMessages.find(
165
+ (m) => m._id === targetMessageId,
166
+ );
167
+ if (targetMessage) {
168
+ text = targetMessage.text;
169
+ } else {
170
+ const targetSearchFields = await ctx.runQuery(
171
+ component.messages.getMessageSearchFields,
172
+ {
173
+ messageId: targetMessageId,
174
+ },
175
+ );
176
+ text = targetSearchFields.text;
177
+ embedding = targetSearchFields.embedding;
178
+ embeddingModel = targetSearchFields.embeddingModel;
179
+ }
180
+ assert(text, "Target message has no text for searching");
181
+ } else if (args.messages?.length) {
182
+ text = extractText(args.messages.at(-1)!);
183
+ assert(text, "Final context message has no text to search");
184
+ }
185
+ assert(text, "No text to search");
98
186
  }
99
- const embeddingFields = opts.searchOptions?.vectorSearch
100
- ? await args.getEmbedding?.(text)
101
- : undefined;
102
- const searchMessages = await ctx.runAction(
187
+ if (opts.searchOptions?.vectorSearch) {
188
+ if (!embedding && args.getEmbedding) {
189
+ const embeddingFields = await args.getEmbedding(text);
190
+ embedding = embeddingFields.embedding;
191
+ const effectiveModel =
192
+ embeddingFields.embeddingModel ?? embeddingFields.textEmbeddingModel;
193
+ embeddingModel = effectiveModel
194
+ ? getModelName(effectiveModel)
195
+ : undefined;
196
+ // TODO: if the text matches the target message, save the embedding
197
+ // for the target message and return the embeddingId on the message.
198
+ }
199
+ }
200
+ const searchResults = await ctx.runAction(
103
201
  component.messages.searchMessages,
104
202
  {
105
203
  searchAllMessagesForUserId: opts?.searchOtherThreads
106
- ? args.userId ??
204
+ ? (args.userId ??
107
205
  (args.threadId &&
108
206
  (
109
207
  await ctx.runQuery(component.threads.getThread, {
110
208
  threadId: args.threadId,
111
209
  })
112
- )?.userId)
210
+ )?.userId))
113
211
  : undefined,
114
212
  threadId: args.threadId,
115
- beforeMessageId: args.upToAndIncludingMessageId,
213
+ targetMessageId,
116
214
  limit: opts.searchOptions?.limit ?? 10,
117
215
  messageRange: {
118
216
  ...DEFAULT_MESSAGE_RANGE,
119
217
  ...opts.searchOptions?.messageRange,
120
218
  },
121
219
  text,
220
+ textSearch: opts.searchOptions?.textSearch,
221
+ vectorSearch: opts.searchOptions?.vectorSearch,
122
222
  vectorScoreThreshold:
123
223
  opts.searchOptions?.vectorScoreThreshold ??
124
224
  DEFAULT_VECTOR_SCORE_THRESHOLD,
125
- embedding: embeddingFields?.embedding,
126
- embeddingModel: embeddingFields?.embeddingModel,
225
+ embedding,
226
+ embeddingModel,
127
227
  },
128
228
  );
129
229
  // TODO: track what messages we used for context
130
- contextMessages.unshift(
131
- ...searchMessages.filter((m) => !included?.has(m._id)),
230
+ searchMessages = filterOutOrphanedToolMessages(
231
+ sorted(searchResults.filter((m) => !included?.has(m._id))),
132
232
  );
133
233
  }
134
234
  // Ensure we don't include tool messages without a corresponding tool call
135
- return filterOutOrphanedToolMessages(
136
- contextMessages.sort((a, b) =>
137
- // Sort the raw MessageDocs by order and stepOrder
138
- a.order === b.order ? a.stepOrder - b.stepOrder : a.order - b.order,
139
- ),
140
- );
235
+ return { recentMessages, searchMessages };
141
236
  }
142
237
 
143
238
  /**
144
239
  * Filter out tool messages that don't have both a tool call and response.
240
+ * For the approval workflow, tool calls with approval responses (but no tool-results yet)
241
+ * should also be kept.
145
242
  * @param docs The messages to filter.
146
243
  * @returns The filtered messages.
147
244
  */
148
245
  export function filterOutOrphanedToolMessages(docs: MessageDoc[]) {
149
246
  const toolCallIds = new Set<string>();
247
+ const toolResultIds = new Set<string>();
248
+ // Track approval workflow: toolCallId → approvalId
249
+ const approvalRequestsByToolCallId = new Map<string, string>();
250
+ // Track which approvalIds have responses
251
+ const approvalResponseIds = new Set<string>();
252
+
150
253
  const result: MessageDoc[] = [];
151
254
  for (const doc of docs) {
152
- if (
153
- doc.message?.role === "assistant" &&
154
- Array.isArray(doc.message.content)
155
- ) {
255
+ if (doc.message && Array.isArray(doc.message.content)) {
156
256
  for (const content of doc.message.content) {
157
257
  if (content.type === "tool-call") {
158
258
  toolCallIds.add(content.toolCallId);
259
+ } else if (content.type === "tool-result") {
260
+ toolResultIds.add(content.toolCallId);
261
+ } else if (content.type === "tool-approval-request") {
262
+ const approvalRequest = content as {
263
+ type: "tool-approval-request";
264
+ toolCallId: string;
265
+ approvalId: string;
266
+ };
267
+ approvalRequestsByToolCallId.set(
268
+ approvalRequest.toolCallId,
269
+ approvalRequest.approvalId,
270
+ );
271
+ } else if (content.type === "tool-approval-response") {
272
+ const approvalResponse = content as {
273
+ type: "tool-approval-response";
274
+ approvalId: string;
275
+ };
276
+ approvalResponseIds.add(approvalResponse.approvalId);
159
277
  }
160
278
  }
161
- result.push(doc);
279
+ }
280
+ }
281
+
282
+ // Helper: check if tool call has a corresponding approval response
283
+ const hasApprovalResponse = (toolCallId: string) => {
284
+ const approvalId = approvalRequestsByToolCallId.get(toolCallId);
285
+ return approvalId !== undefined && approvalResponseIds.has(approvalId);
286
+ };
287
+
288
+ for (const doc of docs) {
289
+ if (
290
+ doc.message?.role === "assistant" &&
291
+ Array.isArray(doc.message.content)
292
+ ) {
293
+ const content = doc.message.content.filter(
294
+ (p) =>
295
+ p.type !== "tool-call" ||
296
+ toolResultIds.has(p.toolCallId) ||
297
+ hasApprovalResponse(p.toolCallId),
298
+ );
299
+ if (content.length) {
300
+ result.push({
301
+ ...doc,
302
+ message: {
303
+ ...doc.message,
304
+ content,
305
+ },
306
+ });
307
+ }
162
308
  } else if (doc.message?.role === "tool") {
163
- if (doc.message.content.every((c) => toolCallIds.has(c.toolCallId))) {
164
- result.push(doc);
165
- } else {
166
- console.debug("Filtering out orphaned tool message", doc);
309
+ const content = doc.message.content.filter((c) => {
310
+ // tool-result parts have toolCallId
311
+ if (c.type === "tool-result") {
312
+ return toolCallIds.has(c.toolCallId);
313
+ }
314
+ // tool-approval-response parts don't have toolCallId, so include them
315
+ return true;
316
+ });
317
+ if (content.length) {
318
+ result.push({
319
+ ...doc,
320
+ message: {
321
+ ...doc.message,
322
+ content,
323
+ },
324
+ });
167
325
  }
168
326
  } else {
169
327
  result.push(doc);
@@ -171,3 +329,352 @@ export function filterOutOrphanedToolMessages(docs: MessageDoc[]) {
171
329
  }
172
330
  return result;
173
331
  }
332
+
333
+ /**
334
+ * Embed a list of messages, including calling any usage handler.
335
+ * This will not save the embeddings to the database.
336
+ */
337
+ export async function embedMessages(
338
+ ctx: ActionCtx,
339
+ {
340
+ userId,
341
+ threadId,
342
+ ...options
343
+ }: {
344
+ userId: string | undefined;
345
+ threadId: string | undefined;
346
+ agentName?: string;
347
+ } & Pick<
348
+ Config,
349
+ "usageHandler" | "textEmbeddingModel" | "embeddingModel" | "callSettings"
350
+ >,
351
+ messages: (ModelMessage | Message)[],
352
+ ): Promise<
353
+ | {
354
+ vectors: (number[] | null)[];
355
+ dimension: VectorDimension;
356
+ model: string;
357
+ }
358
+ | undefined
359
+ > {
360
+ const textEmbeddingModel =
361
+ options.embeddingModel ?? options.textEmbeddingModel;
362
+ if (!textEmbeddingModel) {
363
+ return undefined;
364
+ }
365
+ let embeddings:
366
+ | {
367
+ vectors: (number[] | null)[];
368
+ dimension: VectorDimension;
369
+ model: string;
370
+ }
371
+ | undefined;
372
+ const messageTexts = messages.map((m) => !isTool(m) && extractText(m));
373
+ // Find the indexes of the messages that have text.
374
+ const textIndexes = messageTexts
375
+ .map((t, i) => (t ? i : undefined))
376
+ .filter((i) => i !== undefined);
377
+ if (textIndexes.length === 0) {
378
+ return undefined;
379
+ }
380
+ const values = messageTexts
381
+ .map((t) => t && t.trim().slice(0, MAX_EMBEDDING_TEXT_LENGTH))
382
+ .filter((t): t is string => !!t);
383
+ // Then embed those messages.
384
+ const textEmbeddings = await embedMany(ctx, {
385
+ ...options,
386
+ userId,
387
+ threadId,
388
+ values,
389
+ });
390
+ // Then assemble the embeddings into a single array with nulls for the messages without text.
391
+ const embeddingsOrNull = Array(messages.length).fill(null);
392
+ textIndexes.forEach((i, j) => {
393
+ embeddingsOrNull[i] = textEmbeddings.embeddings[j];
394
+ });
395
+ if (textEmbeddings.embeddings.length > 0) {
396
+ const dimension = textEmbeddings.embeddings[0].length;
397
+ validateVectorDimension(dimension);
398
+ const model = getModelName(textEmbeddingModel);
399
+ embeddings = { vectors: embeddingsOrNull, dimension, model };
400
+ }
401
+ return embeddings;
402
+ }
403
+
404
+ /**
405
+ * Embeds many strings, calling any usage handler.
406
+ * @param ctx The ctx parameter to an action.
407
+ * @param args Arguments to AI SDK's embedMany, and context for the embedding,
408
+ * passed to the usage handler.
409
+ * @returns The embeddings for the strings, matching the order of the values.
410
+ */
411
+ export async function embedMany(
412
+ ctx: ActionCtx,
413
+ args: {
414
+ userId: string | undefined;
415
+ threadId: string | undefined;
416
+ values: string[];
417
+ abortSignal?: AbortSignal;
418
+ headers?: Record<string, string>;
419
+ agentName?: string;
420
+ } & Pick<
421
+ Config,
422
+ "usageHandler" | "textEmbeddingModel" | "embeddingModel" | "callSettings"
423
+ >,
424
+ ): Promise<{ embeddings: number[][] }> {
425
+ const {
426
+ userId,
427
+ threadId,
428
+ values,
429
+ abortSignal,
430
+ headers,
431
+ agentName,
432
+ usageHandler,
433
+ textEmbeddingModel,
434
+ embeddingModel,
435
+ callSettings,
436
+ } = args;
437
+ const effectiveEmbeddingModel = embeddingModel ?? textEmbeddingModel;
438
+ assert(
439
+ effectiveEmbeddingModel,
440
+ "an embeddingModel (or textEmbeddingModel) is required to be set for vector search",
441
+ );
442
+ const result = await embedMany_({
443
+ ...callSettings,
444
+ model: effectiveEmbeddingModel,
445
+ values,
446
+ abortSignal,
447
+ headers,
448
+ });
449
+ if (usageHandler && result.usage) {
450
+ await usageHandler(ctx, {
451
+ userId,
452
+ threadId,
453
+ agentName,
454
+ model: getModelName(effectiveEmbeddingModel),
455
+ provider: getProviderName(effectiveEmbeddingModel),
456
+ providerMetadata: undefined,
457
+ usage: {
458
+ inputTokens: result.usage.tokens,
459
+ outputTokens: 0,
460
+ totalTokens: result.usage.tokens,
461
+ // These detail fields are required by LanguageModelUsage type but we don't
462
+ // have the granular data, so we provide objects with undefined values.
463
+ inputTokenDetails: {
464
+ cacheReadTokens: undefined,
465
+ cacheWriteTokens: undefined,
466
+ noCacheTokens: undefined,
467
+ },
468
+ outputTokenDetails: {
469
+ textTokens: undefined,
470
+ reasoningTokens: undefined,
471
+ },
472
+ },
473
+ });
474
+ }
475
+ return { embeddings: result.embeddings };
476
+ }
477
+
478
+ /**
479
+ * Embed a list of messages, and save the embeddings to the database.
480
+ * @param ctx The ctx parameter to an action.
481
+ * @param component The agent component, usually components.agent.
482
+ * @param args The context for the embedding, passed to the usage handler.
483
+ * @param messages The messages to embed, in the Agent MessageDoc format.
484
+ */
485
+ export async function generateAndSaveEmbeddings(
486
+ ctx: ActionCtx,
487
+ component: AgentComponent,
488
+ args: {
489
+ threadId: string | undefined;
490
+ userId: string | undefined;
491
+ agentName?: string;
492
+ /**
493
+ * @deprecated Use embeddingModel instead.
494
+ */
495
+ textEmbeddingModel?: EmbeddingModel;
496
+ embeddingModel?: EmbeddingModel;
497
+ } & Pick<Config, "usageHandler" | "callSettings">,
498
+ messages: MessageDoc[],
499
+ ) {
500
+ const effectiveEmbeddingModel =
501
+ args.embeddingModel ?? args.textEmbeddingModel;
502
+ if (!effectiveEmbeddingModel) {
503
+ throw new Error(
504
+ "an embeddingModel (or textEmbeddingModel) is required to generate and save embeddings",
505
+ );
506
+ }
507
+ const toEmbed = messages.filter((m) => !m.embeddingId && m.message);
508
+ if (toEmbed.length === 0) {
509
+ return;
510
+ }
511
+ const embeddings = await embedMessages(
512
+ ctx,
513
+ { ...args, embeddingModel: effectiveEmbeddingModel },
514
+ toEmbed.map((m) => m.message!),
515
+ );
516
+ if (embeddings && embeddings.vectors.some((v) => v !== null)) {
517
+ await ctx.runMutation(component.vector.index.insertBatch, {
518
+ vectorDimension: embeddings.dimension,
519
+ vectors: toEmbed
520
+ .map((m, i) => ({
521
+ messageId: m._id,
522
+ model: embeddings.model,
523
+ table: "messages",
524
+ userId: m.userId,
525
+ threadId: m.threadId,
526
+ vector: embeddings.vectors[i]!,
527
+ }))
528
+ .filter((v) => v.vector !== null),
529
+ });
530
+ }
531
+ }
532
+
533
+ /**
534
+ * Similar to fetchContextMessages, but also combines the input messages,
535
+ * with search context, recent messages, input messages, then prompt messages.
536
+ * If there is a promptMessageId and prompt message(s) provided, it will splice
537
+ * the prompt messages into the history to replace the promptMessageId message,
538
+ * but still be followed by any existing messages that were in response to the
539
+ * promptMessageId message.
540
+ */
541
+ export async function fetchContextWithPrompt(
542
+ ctx: ActionCtx,
543
+ component: AgentComponent,
544
+ args: {
545
+ prompt: string | (ModelMessage | Message)[] | undefined;
546
+ messages: (ModelMessage | Message)[] | undefined;
547
+ promptMessageId: string | undefined;
548
+ userId: string | undefined;
549
+ threadId: string | undefined;
550
+ agentName?: string;
551
+ } & Options &
552
+ Config,
553
+ ): Promise<{
554
+ messages: ModelMessage[];
555
+ order: number | undefined;
556
+ stepOrder: number | undefined;
557
+ }> {
558
+ const { threadId, userId, textEmbeddingModel, embeddingModel } = args;
559
+ const effectiveEmbeddingModel = embeddingModel ?? textEmbeddingModel;
560
+
561
+ const promptArray = getPromptArray(args.prompt);
562
+
563
+ const searchText = promptArray.length
564
+ ? extractText(promptArray.at(-1)!)
565
+ : args.promptMessageId
566
+ ? undefined
567
+ : args.messages?.at(-1)
568
+ ? extractText(args.messages.at(-1)!)
569
+ : undefined;
570
+ // If only a messageId is provided, this will add that message to the end.
571
+ const { recentMessages, searchMessages } = await fetchRecentAndSearchMessages(
572
+ ctx,
573
+ component,
574
+ {
575
+ userId,
576
+ threadId,
577
+ targetMessageId: args.promptMessageId,
578
+ searchText,
579
+ contextOptions: args.contextOptions ?? {},
580
+ getEmbedding: async (text) => {
581
+ assert(
582
+ effectiveEmbeddingModel,
583
+ "An embeddingModel (or textEmbeddingModel) is required to be set on the Agent that you're doing vector search with",
584
+ );
585
+ return {
586
+ embedding: (
587
+ await embedMany(ctx, {
588
+ ...args,
589
+ userId,
590
+ values: [text],
591
+ embeddingModel: effectiveEmbeddingModel,
592
+ })
593
+ ).embeddings[0],
594
+ embeddingModel: effectiveEmbeddingModel,
595
+ };
596
+ },
597
+ },
598
+ );
599
+
600
+ const promptMessageIndex = args.promptMessageId
601
+ ? recentMessages.findIndex((m) => m._id === args.promptMessageId)
602
+ : -1;
603
+ const promptMessage =
604
+ promptMessageIndex !== -1 ? recentMessages[promptMessageIndex] : undefined;
605
+ let prePromptDocs = recentMessages;
606
+ const messages = args.messages ?? [];
607
+ let existingResponseDocs: MessageDoc[] = [];
608
+ if (promptMessage) {
609
+ prePromptDocs = recentMessages.slice(0, promptMessageIndex);
610
+ existingResponseDocs = recentMessages.slice(promptMessageIndex + 1);
611
+ if (promptArray.length === 0) {
612
+ // If they didn't override the prompt, use the existing prompt message.
613
+ if (promptMessage.message) {
614
+ promptArray.push(promptMessage.message);
615
+ }
616
+ }
617
+ if (!promptMessage.embeddingId && effectiveEmbeddingModel) {
618
+ // Lazily generate embeddings for the prompt message, if it doesn't have
619
+ // embeddings yet. This can happen if the message was saved in a mutation
620
+ // where the LLM is not available.
621
+ await generateAndSaveEmbeddings(
622
+ ctx,
623
+ component,
624
+ {
625
+ ...args,
626
+ userId,
627
+ embeddingModel: effectiveEmbeddingModel,
628
+ },
629
+ [promptMessage],
630
+ );
631
+ }
632
+ }
633
+
634
+ const search = docsToModelMessages(searchMessages);
635
+ const recent = docsToModelMessages(prePromptDocs);
636
+ const inputMessages = messages.map(toModelMessage);
637
+ const inputPrompt = promptArray.map(toModelMessage);
638
+ const existingResponses = docsToModelMessages(existingResponseDocs);
639
+
640
+ const allMessages = [
641
+ ...search,
642
+ ...recent,
643
+ ...inputMessages,
644
+ ...inputPrompt,
645
+ ...existingResponses,
646
+ ];
647
+ let processedMessages = args.contextHandler
648
+ ? await args.contextHandler(ctx, {
649
+ allMessages,
650
+ search,
651
+ recent,
652
+ inputMessages,
653
+ inputPrompt,
654
+ existingResponses,
655
+ userId,
656
+ threadId,
657
+ })
658
+ : allMessages;
659
+
660
+ // Process messages to inline localhost files (if not, file urls pointing to localhost will be sent to LLM providers)
661
+ if (process.env.CONVEX_CLOUD_URL?.startsWith("http://127.0.0.1")) {
662
+ processedMessages = await inlineMessagesFiles(processedMessages);
663
+ }
664
+
665
+ return {
666
+ messages: processedMessages,
667
+ order: promptMessage?.order,
668
+ stepOrder: promptMessage?.stepOrder,
669
+ };
670
+ }
671
+
672
+ export function getPromptArray(
673
+ prompt: string | (ModelMessage | Message)[] | undefined,
674
+ ): (ModelMessage | Message)[] {
675
+ return !prompt
676
+ ? []
677
+ : Array.isArray(prompt)
678
+ ? prompt
679
+ : [{ role: "user", content: prompt }];
680
+ }