@convex-dev/agent 0.2.6-alpha.0 → 0.2.6
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.
- package/dist/client/definePlaygroundAPI.d.ts +6 -4
- package/dist/client/definePlaygroundAPI.d.ts.map +1 -1
- package/dist/client/definePlaygroundAPI.js +15 -6
- package/dist/client/definePlaygroundAPI.js.map +1 -1
- package/dist/client/index.d.ts +26 -120
- package/dist/client/index.d.ts.map +1 -1
- package/dist/client/index.js +48 -373
- package/dist/client/index.js.map +1 -1
- package/dist/client/messages.d.ts +1 -1
- package/dist/client/messages.d.ts.map +1 -1
- package/dist/client/mockModel.d.ts +3 -3
- package/dist/client/mockModel.d.ts.map +1 -1
- package/dist/client/mockModel.js +22 -17
- package/dist/client/mockModel.js.map +1 -1
- package/dist/client/saveInputMessages.d.ts +20 -0
- package/dist/client/saveInputMessages.d.ts.map +1 -0
- package/dist/client/saveInputMessages.js +57 -0
- package/dist/client/saveInputMessages.js.map +1 -0
- package/dist/client/search.d.ts +110 -9
- package/dist/client/search.d.ts.map +1 -1
- package/dist/client/search.js +271 -39
- package/dist/client/search.js.map +1 -1
- package/dist/client/start.d.ts +83 -0
- package/dist/client/start.d.ts.map +1 -0
- package/dist/client/start.js +171 -0
- package/dist/client/start.js.map +1 -0
- package/dist/client/streaming.d.ts +8 -8
- package/dist/client/streaming.d.ts.map +1 -1
- package/dist/client/streaming.js +2 -1
- package/dist/client/streaming.js.map +1 -1
- package/dist/client/textStreamParts.d.ts.map +1 -1
- package/dist/client/textStreamParts.js +2 -9
- package/dist/client/textStreamParts.js.map +1 -1
- package/dist/client/threads.d.ts +1 -1
- package/dist/client/threads.d.ts.map +1 -1
- package/dist/client/types.d.ts +137 -5
- package/dist/client/types.d.ts.map +1 -1
- package/dist/component/_generated/api.d.ts +11 -3
- package/dist/component/messages.d.ts +13 -4
- package/dist/component/messages.d.ts.map +1 -1
- package/dist/component/messages.js +67 -25
- package/dist/component/messages.js.map +1 -1
- package/dist/component/schema.d.ts +2 -1643
- package/dist/component/schema.d.ts.map +1 -1
- package/dist/component/schema.js +0 -24
- package/dist/component/schema.js.map +1 -1
- package/dist/mapping.d.ts +7 -9
- package/dist/mapping.d.ts.map +1 -1
- package/dist/mapping.js +73 -7
- package/dist/mapping.js.map +1 -1
- package/dist/react/deltas.d.ts.map +1 -1
- package/dist/react/deltas.js +15 -5
- package/dist/react/deltas.js.map +1 -1
- package/dist/react/fromUIMessages.d.ts +13 -0
- package/dist/react/fromUIMessages.d.ts.map +1 -0
- package/dist/react/fromUIMessages.js +70 -0
- package/dist/react/fromUIMessages.js.map +1 -0
- package/dist/react/toUIMessages.d.ts +5 -2
- package/dist/react/toUIMessages.d.ts.map +1 -1
- package/dist/react/toUIMessages.js +3 -0
- package/dist/react/toUIMessages.js.map +1 -1
- package/dist/shared.d.ts +10 -0
- package/dist/shared.d.ts.map +1 -1
- package/dist/shared.js +26 -0
- package/dist/shared.js.map +1 -1
- package/dist/validators.d.ts +1640 -0
- package/dist/validators.d.ts.map +1 -1
- package/dist/validators.js +41 -0
- package/dist/validators.js.map +1 -1
- package/package.json +1 -1
- package/src/client/definePlaygroundAPI.ts +16 -7
- package/src/client/index.test.ts +11 -46
- package/src/client/index.ts +99 -558
- package/src/client/messages.ts +1 -1
- package/src/client/mock.json +68 -0
- package/src/client/mockModel.ts +34 -23
- package/src/client/saveInputMessages.test.ts +576 -0
- package/src/client/saveInputMessages.ts +100 -0
- package/src/client/search.test.ts +1017 -0
- package/src/client/search.ts +446 -68
- package/src/client/start.ts +313 -0
- package/src/client/stream.json +48 -0
- package/src/client/streaming.ts +3 -3
- package/src/client/textStreamParts.ts +2 -11
- package/src/client/threads.ts +1 -1
- package/src/client/types.ts +143 -3
- package/src/component/_generated/api.d.ts +11 -3
- package/src/component/messages.ts +73 -27
- package/src/component/schema.ts +1 -29
- package/src/mapping.ts +84 -7
- package/src/node_modules/.vite/vitest/da39a3ee5e6b4b0d3255bfef95601890afd80709/results.json +1 -0
- package/src/react/deltas.ts +18 -5
- package/src/react/fromUIMessages.test.ts +427 -0
- package/src/react/fromUIMessages.ts +85 -0
- package/src/react/toUIMessages.ts +21 -13
- package/src/shared.ts +33 -0
- package/src/validators.test.ts +13 -2
- package/src/validators.ts +48 -0
package/src/client/search.ts
CHANGED
|
@@ -1,25 +1,41 @@
|
|
|
1
|
-
import
|
|
2
|
-
|
|
3
|
-
|
|
4
|
-
|
|
5
|
-
|
|
6
|
-
} from "./types.js";
|
|
7
|
-
import type { MessageDoc } from "../component/schema.js";
|
|
8
|
-
import type { EmbeddingModel, LanguageModel, 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,
|
|
14
19
|
sorted,
|
|
15
20
|
} from "../shared.js";
|
|
16
21
|
import type { Message } from "../validators.js";
|
|
22
|
+
import type {
|
|
23
|
+
AgentComponent,
|
|
24
|
+
Config,
|
|
25
|
+
ContextOptions,
|
|
26
|
+
Options,
|
|
27
|
+
RunActionCtx,
|
|
28
|
+
RunQueryCtx,
|
|
29
|
+
} from "./types.js";
|
|
30
|
+
import { inlineMessagesFiles } from "./files.js";
|
|
31
|
+
import { deserializeMessage } from "../mapping.js";
|
|
17
32
|
|
|
18
33
|
const DEFAULT_VECTOR_SCORE_THRESHOLD = 0.0;
|
|
34
|
+
// 10k characters should be more than enough for most cases, and stays under
|
|
35
|
+
// the 8k token limit for some models.
|
|
36
|
+
const MAX_EMBEDDING_TEXT_LENGTH = 10_000;
|
|
19
37
|
|
|
20
|
-
export type GetEmbedding = (
|
|
21
|
-
text: string,
|
|
22
|
-
) => Promise<{
|
|
38
|
+
export type GetEmbedding = (text: string) => Promise<{
|
|
23
39
|
embedding: number[];
|
|
24
40
|
textEmbeddingModel: string | EmbeddingModel<string>;
|
|
25
41
|
}>;
|
|
@@ -38,26 +54,73 @@ export async function fetchContextMessages(
|
|
|
38
54
|
args: {
|
|
39
55
|
userId: string | undefined;
|
|
40
56
|
threadId: string | undefined;
|
|
41
|
-
messages: (ModelMessage | Message)[];
|
|
42
57
|
/**
|
|
43
|
-
* If
|
|
44
|
-
*
|
|
45
|
-
|
|
58
|
+
* If targetMessageId is not provided, this text will be used
|
|
59
|
+
* for text and vector search
|
|
60
|
+
*/
|
|
61
|
+
searchText?: string;
|
|
62
|
+
/**
|
|
63
|
+
* If provided, it will use this message for text/vector search (if enabled)
|
|
64
|
+
* and will only fetch messages up to (and including) this message's "order"
|
|
65
|
+
*/
|
|
66
|
+
targetMessageId?: string;
|
|
67
|
+
/**
|
|
68
|
+
* @deprecated use searchText and targetMessageId instead
|
|
69
|
+
*/
|
|
70
|
+
messages?: (ModelMessage | Message)[];
|
|
71
|
+
/**
|
|
72
|
+
* @deprecated use targetMessageId instead
|
|
46
73
|
*/
|
|
47
74
|
upToAndIncludingMessageId?: string;
|
|
48
75
|
contextOptions: ContextOptions;
|
|
49
76
|
getEmbedding?: GetEmbedding;
|
|
50
77
|
},
|
|
51
78
|
): Promise<MessageDoc[]> {
|
|
79
|
+
const { recentMessages, searchMessages } = await fetchRecentAndSearchMessages(
|
|
80
|
+
ctx,
|
|
81
|
+
component,
|
|
82
|
+
args,
|
|
83
|
+
);
|
|
84
|
+
return [...searchMessages, ...recentMessages];
|
|
85
|
+
}
|
|
86
|
+
|
|
87
|
+
export async function fetchRecentAndSearchMessages(
|
|
88
|
+
ctx: RunQueryCtx | RunActionCtx,
|
|
89
|
+
component: AgentComponent,
|
|
90
|
+
args: {
|
|
91
|
+
userId: string | undefined;
|
|
92
|
+
threadId: string | undefined;
|
|
93
|
+
/**
|
|
94
|
+
* If targetMessageId is not provided, this text will be used
|
|
95
|
+
* for text and vector search
|
|
96
|
+
*/
|
|
97
|
+
searchText?: string;
|
|
98
|
+
/**
|
|
99
|
+
* If provided, it will use this message for text/vector search (if enabled)
|
|
100
|
+
* and will only fetch messages up to (and including) this message's "order"
|
|
101
|
+
*/
|
|
102
|
+
targetMessageId?: string;
|
|
103
|
+
/**
|
|
104
|
+
* @deprecated use searchText and targetMessageId instead
|
|
105
|
+
*/
|
|
106
|
+
messages?: (ModelMessage | Message)[];
|
|
107
|
+
/**
|
|
108
|
+
* @deprecated use targetMessageId instead
|
|
109
|
+
*/
|
|
110
|
+
upToAndIncludingMessageId?: string;
|
|
111
|
+
contextOptions: ContextOptions;
|
|
112
|
+
getEmbedding?: GetEmbedding;
|
|
113
|
+
},
|
|
114
|
+
): Promise<{ recentMessages: MessageDoc[]; searchMessages: MessageDoc[] }> {
|
|
52
115
|
assert(args.userId || args.threadId, "Specify userId or threadId");
|
|
53
116
|
const opts = args.contextOptions;
|
|
54
117
|
// Fetch the latest messages from the thread
|
|
55
118
|
let included: Set<string> | undefined;
|
|
56
|
-
|
|
57
|
-
|
|
58
|
-
|
|
59
|
-
|
|
60
|
-
) {
|
|
119
|
+
let recentMessages: MessageDoc[] = [];
|
|
120
|
+
let searchMessages: MessageDoc[] = [];
|
|
121
|
+
const targetMessageId =
|
|
122
|
+
args.targetMessageId ?? args.upToAndIncludingMessageId;
|
|
123
|
+
if (args.threadId && opts.recentMessages !== 0) {
|
|
61
124
|
const { page } = await ctx.runQuery(
|
|
62
125
|
component.messages.listMessagesByThreadId,
|
|
63
126
|
{
|
|
@@ -67,46 +130,59 @@ export async function fetchContextMessages(
|
|
|
67
130
|
numItems: opts.recentMessages ?? DEFAULT_RECENT_MESSAGES,
|
|
68
131
|
cursor: null,
|
|
69
132
|
},
|
|
70
|
-
upToAndIncludingMessageId:
|
|
133
|
+
upToAndIncludingMessageId: targetMessageId,
|
|
71
134
|
order: "desc",
|
|
72
135
|
statuses: ["success"],
|
|
73
136
|
},
|
|
74
137
|
);
|
|
75
138
|
included = new Set(page.map((m) => m._id));
|
|
76
|
-
|
|
77
|
-
// Reverse since we fetched in descending order
|
|
78
|
-
...page.reverse(),
|
|
79
|
-
);
|
|
139
|
+
recentMessages = filterOutOrphanedToolMessages(sorted(page));
|
|
80
140
|
}
|
|
81
141
|
if (
|
|
82
142
|
(opts.searchOptions?.textSearch || opts.searchOptions?.vectorSearch) &&
|
|
83
143
|
opts.searchOptions?.limit
|
|
84
144
|
) {
|
|
85
|
-
const targetMessage = contextMessages.find(
|
|
86
|
-
(m) => m._id === args.upToAndIncludingMessageId,
|
|
87
|
-
)?.message;
|
|
88
|
-
const messagesToSearch = targetMessage ? [targetMessage] : args.messages;
|
|
89
145
|
if (!("runAction" in ctx)) {
|
|
90
146
|
throw new Error("searchUserMessages only works in an action");
|
|
91
147
|
}
|
|
92
|
-
|
|
93
|
-
|
|
94
|
-
|
|
95
|
-
|
|
96
|
-
|
|
97
|
-
|
|
98
|
-
|
|
99
|
-
|
|
100
|
-
|
|
101
|
-
|
|
102
|
-
|
|
103
|
-
|
|
148
|
+
let text = args.searchText;
|
|
149
|
+
let embedding: number[] | undefined;
|
|
150
|
+
let embeddingModel: string | undefined;
|
|
151
|
+
if (!text) {
|
|
152
|
+
if (targetMessageId) {
|
|
153
|
+
const targetMessage = recentMessages.find(
|
|
154
|
+
(m) => m._id === targetMessageId,
|
|
155
|
+
);
|
|
156
|
+
if (targetMessage) {
|
|
157
|
+
text = targetMessage.text;
|
|
158
|
+
} else {
|
|
159
|
+
const targetSearchFields = await ctx.runQuery(
|
|
160
|
+
component.messages.getMessageSearchFields,
|
|
161
|
+
{
|
|
162
|
+
messageId: targetMessageId,
|
|
163
|
+
},
|
|
164
|
+
);
|
|
165
|
+
text = targetSearchFields.text;
|
|
166
|
+
embedding = targetSearchFields.embedding;
|
|
167
|
+
embeddingModel = targetSearchFields.embeddingModel;
|
|
168
|
+
}
|
|
169
|
+
assert(text, "Target message has no text for searching");
|
|
170
|
+
} else if (args.messages?.length) {
|
|
171
|
+
text = extractText(args.messages.at(-1)!);
|
|
172
|
+
assert(text, "Final context message has no text to search");
|
|
173
|
+
}
|
|
174
|
+
assert(text, "No text to search");
|
|
104
175
|
}
|
|
105
|
-
|
|
106
|
-
|
|
107
|
-
|
|
108
|
-
|
|
109
|
-
|
|
176
|
+
if (opts.searchOptions?.vectorSearch) {
|
|
177
|
+
if (!embedding && args.getEmbedding) {
|
|
178
|
+
const embeddingFields = await args.getEmbedding(text);
|
|
179
|
+
embedding = embeddingFields.embedding;
|
|
180
|
+
embeddingModel = embeddingFields.textEmbeddingModel
|
|
181
|
+
? getModelName(embeddingFields.textEmbeddingModel)
|
|
182
|
+
: undefined;
|
|
183
|
+
}
|
|
184
|
+
}
|
|
185
|
+
const searchResults = await ctx.runAction(
|
|
110
186
|
component.messages.searchMessages,
|
|
111
187
|
{
|
|
112
188
|
searchAllMessagesForUserId: opts?.searchOtherThreads
|
|
@@ -119,29 +195,29 @@ export async function fetchContextMessages(
|
|
|
119
195
|
)?.userId))
|
|
120
196
|
: undefined,
|
|
121
197
|
threadId: args.threadId,
|
|
122
|
-
|
|
198
|
+
targetMessageId,
|
|
123
199
|
limit: opts.searchOptions?.limit ?? 10,
|
|
124
200
|
messageRange: {
|
|
125
201
|
...DEFAULT_MESSAGE_RANGE,
|
|
126
202
|
...opts.searchOptions?.messageRange,
|
|
127
203
|
},
|
|
128
204
|
text,
|
|
205
|
+
textSearch: opts.searchOptions?.textSearch,
|
|
206
|
+
vectorSearch: opts.searchOptions?.vectorSearch,
|
|
129
207
|
vectorScoreThreshold:
|
|
130
208
|
opts.searchOptions?.vectorScoreThreshold ??
|
|
131
209
|
DEFAULT_VECTOR_SCORE_THRESHOLD,
|
|
132
|
-
embedding
|
|
133
|
-
embeddingModel
|
|
134
|
-
? getModelName(embeddingFields.textEmbeddingModel)
|
|
135
|
-
: undefined,
|
|
210
|
+
embedding,
|
|
211
|
+
embeddingModel,
|
|
136
212
|
},
|
|
137
213
|
);
|
|
138
214
|
// TODO: track what messages we used for context
|
|
139
|
-
|
|
140
|
-
|
|
215
|
+
searchMessages = filterOutOrphanedToolMessages(
|
|
216
|
+
sorted(searchResults.filter((m) => !included?.has(m._id))),
|
|
141
217
|
);
|
|
142
218
|
}
|
|
143
219
|
// Ensure we don't include tool messages without a corresponding tool call
|
|
144
|
-
return
|
|
220
|
+
return { recentMessages, searchMessages };
|
|
145
221
|
}
|
|
146
222
|
|
|
147
223
|
/**
|
|
@@ -176,23 +252,325 @@ export function filterOutOrphanedToolMessages(docs: MessageDoc[]) {
|
|
|
176
252
|
return result;
|
|
177
253
|
}
|
|
178
254
|
|
|
179
|
-
|
|
180
|
-
|
|
181
|
-
|
|
182
|
-
|
|
183
|
-
|
|
184
|
-
|
|
255
|
+
/**
|
|
256
|
+
* Embed a list of messages, including calling any usage handler.
|
|
257
|
+
* This will not save the embeddings to the database.
|
|
258
|
+
*/
|
|
259
|
+
export async function embedMessages(
|
|
260
|
+
ctx: RunActionCtx,
|
|
261
|
+
{
|
|
262
|
+
userId,
|
|
263
|
+
threadId,
|
|
264
|
+
...options
|
|
265
|
+
}: {
|
|
266
|
+
userId: string | undefined;
|
|
267
|
+
threadId: string | undefined;
|
|
268
|
+
agentName?: string;
|
|
269
|
+
} & Pick<Config, "usageHandler" | "textEmbeddingModel" | "callSettings">,
|
|
270
|
+
messages: (ModelMessage | Message)[],
|
|
271
|
+
): Promise<
|
|
272
|
+
| {
|
|
273
|
+
vectors: (number[] | null)[];
|
|
274
|
+
dimension: VectorDimension;
|
|
275
|
+
model: string;
|
|
185
276
|
}
|
|
186
|
-
|
|
277
|
+
| undefined
|
|
278
|
+
> {
|
|
279
|
+
if (!options.textEmbeddingModel) {
|
|
280
|
+
return undefined;
|
|
281
|
+
}
|
|
282
|
+
let embeddings:
|
|
283
|
+
| {
|
|
284
|
+
vectors: (number[] | null)[];
|
|
285
|
+
dimension: VectorDimension;
|
|
286
|
+
model: string;
|
|
287
|
+
}
|
|
288
|
+
| undefined;
|
|
289
|
+
const messageTexts = messages.map((m) => !isTool(m) && extractText(m));
|
|
290
|
+
// Find the indexes of the messages that have text.
|
|
291
|
+
const textIndexes = messageTexts
|
|
292
|
+
.map((t, i) => (t ? i : undefined))
|
|
293
|
+
.filter((i) => i !== undefined);
|
|
294
|
+
if (textIndexes.length === 0) {
|
|
295
|
+
return undefined;
|
|
296
|
+
}
|
|
297
|
+
const values = messageTexts
|
|
298
|
+
.map((t) => t && t.trim().slice(0, MAX_EMBEDDING_TEXT_LENGTH))
|
|
299
|
+
.filter((t): t is string => !!t);
|
|
300
|
+
// Then embed those messages.
|
|
301
|
+
const textEmbeddings = await embedMany(ctx, {
|
|
302
|
+
...options,
|
|
303
|
+
userId,
|
|
304
|
+
threadId,
|
|
305
|
+
values,
|
|
306
|
+
});
|
|
307
|
+
// Then assemble the embeddings into a single array with nulls for the messages without text.
|
|
308
|
+
const embeddingsOrNull = Array(messages.length).fill(null);
|
|
309
|
+
textIndexes.forEach((i, j) => {
|
|
310
|
+
embeddingsOrNull[i] = textEmbeddings.embeddings[j];
|
|
311
|
+
});
|
|
312
|
+
if (textEmbeddings.embeddings.length > 0) {
|
|
313
|
+
const dimension = textEmbeddings.embeddings[0].length;
|
|
314
|
+
validateVectorDimension(dimension);
|
|
315
|
+
const model = getModelName(options.textEmbeddingModel);
|
|
316
|
+
embeddings = { vectors: embeddingsOrNull, dimension, model };
|
|
317
|
+
}
|
|
318
|
+
return embeddings;
|
|
319
|
+
}
|
|
320
|
+
|
|
321
|
+
/**
|
|
322
|
+
* Embeds many strings, calling any usage handler.
|
|
323
|
+
* @param ctx The ctx parameter to an action.
|
|
324
|
+
* @param args Arguments to AI SDK's embedMany, and context for the embedding,
|
|
325
|
+
* passed to the usage handler.
|
|
326
|
+
* @returns The embeddings for the strings, matching the order of the values.
|
|
327
|
+
*/
|
|
328
|
+
export async function embedMany(
|
|
329
|
+
ctx: RunActionCtx,
|
|
330
|
+
{
|
|
331
|
+
userId,
|
|
332
|
+
threadId,
|
|
333
|
+
values,
|
|
334
|
+
abortSignal,
|
|
335
|
+
headers,
|
|
336
|
+
agentName,
|
|
337
|
+
usageHandler,
|
|
338
|
+
textEmbeddingModel,
|
|
339
|
+
callSettings,
|
|
340
|
+
}: {
|
|
341
|
+
userId: string | undefined;
|
|
342
|
+
threadId: string | undefined;
|
|
343
|
+
values: string[];
|
|
344
|
+
abortSignal?: AbortSignal;
|
|
345
|
+
headers?: Record<string, string>;
|
|
346
|
+
agentName?: string;
|
|
347
|
+
} & Pick<Config, "usageHandler" | "textEmbeddingModel" | "callSettings">,
|
|
348
|
+
): Promise<{ embeddings: number[][] }> {
|
|
349
|
+
const embeddingModel = textEmbeddingModel;
|
|
350
|
+
assert(
|
|
351
|
+
embeddingModel,
|
|
352
|
+
"a textEmbeddingModel is required to be set for vector search",
|
|
353
|
+
);
|
|
354
|
+
const result = await embedMany_({
|
|
355
|
+
...callSettings,
|
|
356
|
+
model: embeddingModel,
|
|
357
|
+
values,
|
|
358
|
+
abortSignal,
|
|
359
|
+
headers,
|
|
360
|
+
});
|
|
361
|
+
if (usageHandler && result.usage) {
|
|
362
|
+
await usageHandler(ctx, {
|
|
363
|
+
userId,
|
|
364
|
+
threadId,
|
|
365
|
+
agentName,
|
|
366
|
+
model: getModelName(embeddingModel),
|
|
367
|
+
provider: getProviderName(embeddingModel),
|
|
368
|
+
providerMetadata: undefined,
|
|
369
|
+
usage: {
|
|
370
|
+
inputTokens: result.usage.tokens,
|
|
371
|
+
outputTokens: 0,
|
|
372
|
+
totalTokens: result.usage.tokens,
|
|
373
|
+
},
|
|
374
|
+
});
|
|
375
|
+
}
|
|
376
|
+
return { embeddings: result.embeddings };
|
|
377
|
+
}
|
|
378
|
+
|
|
379
|
+
/**
|
|
380
|
+
* Embed a list of messages, and save the embeddings to the database.
|
|
381
|
+
* @param ctx The ctx parameter to an action.
|
|
382
|
+
* @param component The agent component, usually components.agent.
|
|
383
|
+
* @param args The context for the embedding, passed to the usage handler.
|
|
384
|
+
* @param messages The messages to embed, in the Agent MessageDoc format.
|
|
385
|
+
*/
|
|
386
|
+
export async function generateAndSaveEmbeddings(
|
|
387
|
+
ctx: RunActionCtx,
|
|
388
|
+
component: AgentComponent,
|
|
389
|
+
args: {
|
|
390
|
+
threadId: string | undefined;
|
|
391
|
+
userId: string | undefined;
|
|
392
|
+
agentName?: string;
|
|
393
|
+
textEmbeddingModel: EmbeddingModel<string>;
|
|
394
|
+
} & Pick<Config, "usageHandler" | "callSettings">,
|
|
395
|
+
messages: MessageDoc[],
|
|
396
|
+
) {
|
|
397
|
+
const toEmbed = messages.filter((m) => !m.embeddingId && m.message);
|
|
398
|
+
if (toEmbed.length === 0) {
|
|
399
|
+
return;
|
|
400
|
+
}
|
|
401
|
+
const embeddings = await embedMessages(
|
|
402
|
+
ctx,
|
|
403
|
+
args,
|
|
404
|
+
toEmbed.map((m) => m.message!),
|
|
405
|
+
);
|
|
406
|
+
if (embeddings && embeddings.vectors.some((v) => v !== null)) {
|
|
407
|
+
await ctx.runMutation(component.vector.index.insertBatch, {
|
|
408
|
+
vectorDimension: embeddings.dimension,
|
|
409
|
+
vectors: toEmbed
|
|
410
|
+
.map((m, i) => ({
|
|
411
|
+
messageId: m._id,
|
|
412
|
+
model: embeddings.model,
|
|
413
|
+
table: "messages",
|
|
414
|
+
userId: m.userId,
|
|
415
|
+
threadId: m.threadId,
|
|
416
|
+
vector: embeddings.vectors[i]!,
|
|
417
|
+
}))
|
|
418
|
+
.filter((v) => v.vector !== null),
|
|
419
|
+
});
|
|
187
420
|
}
|
|
188
|
-
return embeddingModel.modelId;
|
|
189
421
|
}
|
|
190
422
|
|
|
191
|
-
|
|
192
|
-
|
|
193
|
-
|
|
194
|
-
|
|
195
|
-
|
|
423
|
+
/**
|
|
424
|
+
* Similar to fetchContextMessages, but also combines the input messages,
|
|
425
|
+
* with search context, recent messages, input messages, then prompt messages.
|
|
426
|
+
* If there is a promptMessageId and prompt message(s) provided, it will splice
|
|
427
|
+
* the prompt messages into the history to replace the promptMessageId message,
|
|
428
|
+
* but still be followed by any existing messages that were in response to the
|
|
429
|
+
* promptMessageId message.
|
|
430
|
+
*/
|
|
431
|
+
export async function fetchContextWithPrompt(
|
|
432
|
+
ctx: RunActionCtx,
|
|
433
|
+
component: AgentComponent,
|
|
434
|
+
args: {
|
|
435
|
+
prompt: string | (ModelMessage | Message)[] | undefined;
|
|
436
|
+
messages: (ModelMessage | Message)[] | undefined;
|
|
437
|
+
promptMessageId: string | undefined;
|
|
438
|
+
userId: string | undefined;
|
|
439
|
+
threadId: string | undefined;
|
|
440
|
+
agentName?: string;
|
|
441
|
+
} & Options &
|
|
442
|
+
Config,
|
|
443
|
+
): Promise<{
|
|
444
|
+
messages: ModelMessage[];
|
|
445
|
+
order: number | undefined;
|
|
446
|
+
stepOrder: number | undefined;
|
|
447
|
+
}> {
|
|
448
|
+
const { threadId, userId, textEmbeddingModel } = args;
|
|
449
|
+
|
|
450
|
+
const promptArray = getPromptArray(args.prompt);
|
|
451
|
+
|
|
452
|
+
const searchText = promptArray.length
|
|
453
|
+
? extractText(promptArray.at(-1)!)
|
|
454
|
+
: args.promptMessageId
|
|
455
|
+
? undefined
|
|
456
|
+
: args.messages?.at(-1)
|
|
457
|
+
? extractText(args.messages.at(-1)!)
|
|
458
|
+
: undefined;
|
|
459
|
+
// If only a messageId is provided, this will add that message to the end.
|
|
460
|
+
const { recentMessages, searchMessages } = await fetchRecentAndSearchMessages(
|
|
461
|
+
ctx,
|
|
462
|
+
component,
|
|
463
|
+
{
|
|
464
|
+
userId,
|
|
465
|
+
threadId,
|
|
466
|
+
targetMessageId: args.promptMessageId,
|
|
467
|
+
searchText,
|
|
468
|
+
contextOptions: args.contextOptions ?? {},
|
|
469
|
+
getEmbedding: async (text) => {
|
|
470
|
+
assert(
|
|
471
|
+
textEmbeddingModel,
|
|
472
|
+
"A textEmbeddingModel is required to be set on the Agent that you're doing vector search with",
|
|
473
|
+
);
|
|
474
|
+
return {
|
|
475
|
+
embedding: (
|
|
476
|
+
await embedMany(ctx, {
|
|
477
|
+
...args,
|
|
478
|
+
userId,
|
|
479
|
+
values: [text],
|
|
480
|
+
textEmbeddingModel,
|
|
481
|
+
})
|
|
482
|
+
).embeddings[0],
|
|
483
|
+
textEmbeddingModel,
|
|
484
|
+
};
|
|
485
|
+
},
|
|
486
|
+
},
|
|
487
|
+
);
|
|
488
|
+
|
|
489
|
+
const promptMessageIndex = args.promptMessageId
|
|
490
|
+
? recentMessages.findIndex((m) => m._id === args.promptMessageId)
|
|
491
|
+
: -1;
|
|
492
|
+
const promptMessage =
|
|
493
|
+
promptMessageIndex !== -1 ? recentMessages[promptMessageIndex] : undefined;
|
|
494
|
+
let prePromptDocs = recentMessages;
|
|
495
|
+
const messages = args.messages ?? [];
|
|
496
|
+
let existingResponseDocs: MessageDoc[] = [];
|
|
497
|
+
if (promptMessage) {
|
|
498
|
+
prePromptDocs = recentMessages.slice(0, promptMessageIndex);
|
|
499
|
+
existingResponseDocs = recentMessages.slice(promptMessageIndex + 1);
|
|
500
|
+
if (promptArray.length === 0) {
|
|
501
|
+
// If they didn't override the prompt, use the existing prompt message.
|
|
502
|
+
if (promptMessage.message) {
|
|
503
|
+
promptArray.push(promptMessage.message);
|
|
504
|
+
}
|
|
505
|
+
}
|
|
506
|
+
if (!promptMessage.embeddingId && textEmbeddingModel) {
|
|
507
|
+
// Lazily generate embeddings for the prompt message, if it doesn't have
|
|
508
|
+
// embeddings yet. This can happen if the message was saved in a mutation
|
|
509
|
+
// where the LLM is not available.
|
|
510
|
+
await generateAndSaveEmbeddings(
|
|
511
|
+
ctx,
|
|
512
|
+
component,
|
|
513
|
+
{
|
|
514
|
+
...args,
|
|
515
|
+
userId,
|
|
516
|
+
textEmbeddingModel,
|
|
517
|
+
},
|
|
518
|
+
[promptMessage],
|
|
519
|
+
);
|
|
520
|
+
}
|
|
521
|
+
}
|
|
522
|
+
|
|
523
|
+
const search = searchMessages
|
|
524
|
+
.map((m) => m.message)
|
|
525
|
+
.filter((m) => !!m)
|
|
526
|
+
.map(deserializeMessage);
|
|
527
|
+
const recent = prePromptDocs
|
|
528
|
+
.map((m) => m.message)
|
|
529
|
+
.filter((m) => !!m)
|
|
530
|
+
.map(deserializeMessage);
|
|
531
|
+
const inputMessages = messages.map(deserializeMessage);
|
|
532
|
+
const inputPrompt = promptArray.map(deserializeMessage);
|
|
533
|
+
const existingResponses = existingResponseDocs
|
|
534
|
+
.map((m) => m.message)
|
|
535
|
+
.filter((m) => !!m)
|
|
536
|
+
.map(deserializeMessage);
|
|
537
|
+
|
|
538
|
+
let processedMessages = args.contextHandler
|
|
539
|
+
? await args.contextHandler(ctx, {
|
|
540
|
+
search,
|
|
541
|
+
recent,
|
|
542
|
+
inputMessages,
|
|
543
|
+
inputPrompt,
|
|
544
|
+
existingResponses,
|
|
545
|
+
userId,
|
|
546
|
+
threadId,
|
|
547
|
+
})
|
|
548
|
+
: [
|
|
549
|
+
...search,
|
|
550
|
+
...recent,
|
|
551
|
+
...inputMessages,
|
|
552
|
+
...inputPrompt,
|
|
553
|
+
...existingResponses,
|
|
554
|
+
];
|
|
555
|
+
|
|
556
|
+
// Process messages to inline localhost files (if not, file urls pointing to localhost will be sent to LLM providers)
|
|
557
|
+
if (process.env.CONVEX_CLOUD_URL?.startsWith("http://127.0.0.1")) {
|
|
558
|
+
processedMessages = await inlineMessagesFiles(processedMessages);
|
|
196
559
|
}
|
|
197
|
-
|
|
560
|
+
|
|
561
|
+
return {
|
|
562
|
+
messages: processedMessages,
|
|
563
|
+
order: promptMessage?.order,
|
|
564
|
+
stepOrder: promptMessage?.stepOrder,
|
|
565
|
+
};
|
|
566
|
+
}
|
|
567
|
+
|
|
568
|
+
export function getPromptArray(
|
|
569
|
+
prompt: string | (ModelMessage | Message)[] | undefined,
|
|
570
|
+
): (ModelMessage | Message)[] {
|
|
571
|
+
return !prompt
|
|
572
|
+
? []
|
|
573
|
+
: Array.isArray(prompt)
|
|
574
|
+
? prompt
|
|
575
|
+
: [{ role: "user", content: prompt }];
|
|
198
576
|
}
|