@convex-dev/agent 0.6.0 → 0.6.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.
- package/dist/UIMessages.d.ts +46 -0
- package/dist/UIMessages.d.ts.map +1 -0
- package/dist/UIMessages.js +634 -0
- package/dist/UIMessages.js.map +1 -0
- package/dist/client/createTool.d.ts +167 -0
- package/dist/client/createTool.d.ts.map +1 -0
- package/dist/client/createTool.js +116 -0
- package/dist/client/createTool.js.map +1 -0
- package/dist/client/defaultComponent.d.ts +11 -0
- package/dist/client/defaultComponent.d.ts.map +1 -0
- package/dist/client/defaultComponent.js +7 -0
- package/dist/client/defaultComponent.js.map +1 -0
- package/dist/client/definePlaygroundAPI.d.ts +1725 -0
- package/dist/client/definePlaygroundAPI.d.ts.map +1 -0
- package/dist/client/definePlaygroundAPI.js +271 -0
- package/dist/client/definePlaygroundAPI.js.map +1 -0
- package/dist/client/files.d.ts +69 -0
- package/dist/client/files.d.ts.map +1 -0
- package/dist/client/files.js +181 -0
- package/dist/client/files.js.map +1 -0
- package/dist/client/index.d.ts +2091 -0
- package/dist/client/index.d.ts.map +1 -0
- package/dist/client/index.js +929 -0
- package/dist/client/index.js.map +1 -0
- package/dist/client/messages.d.ts +461 -0
- package/dist/client/messages.d.ts.map +1 -0
- package/dist/client/messages.js +106 -0
- package/dist/client/messages.js.map +1 -0
- package/dist/client/mockModel.d.ts +42 -0
- package/dist/client/mockModel.d.ts.map +1 -0
- package/dist/client/mockModel.js +182 -0
- package/dist/client/mockModel.js.map +1 -0
- package/dist/client/saveInputMessages.d.ts +20 -0
- package/dist/client/saveInputMessages.d.ts.map +1 -0
- package/dist/client/saveInputMessages.js +58 -0
- package/dist/client/saveInputMessages.js.map +1 -0
- package/dist/client/search.d.ts +493 -0
- package/dist/client/search.d.ts.map +1 -0
- package/dist/client/search.js +435 -0
- package/dist/client/search.js.map +1 -0
- package/dist/client/start.d.ts +84 -0
- package/dist/client/start.d.ts.map +1 -0
- package/dist/client/start.js +185 -0
- package/dist/client/start.js.map +1 -0
- package/dist/client/streamText.d.ts +46 -0
- package/dist/client/streamText.d.ts.map +1 -0
- package/dist/client/streamText.js +127 -0
- package/dist/client/streamText.js.map +1 -0
- package/dist/client/streaming.d.ts +3778 -0
- package/dist/client/streaming.d.ts.map +1 -0
- package/dist/client/streaming.js +330 -0
- package/dist/client/streaming.js.map +1 -0
- package/dist/client/threads.d.ts +46 -0
- package/dist/client/threads.d.ts.map +1 -0
- package/dist/client/threads.js +49 -0
- package/dist/client/threads.js.map +1 -0
- package/dist/client/types.d.ts +461 -0
- package/dist/client/types.d.ts.map +1 -0
- package/dist/client/types.js +2 -0
- package/dist/client/types.js.map +1 -0
- package/dist/client/utils.d.ts +4 -0
- package/dist/client/utils.d.ts.map +1 -0
- package/dist/client/utils.js +21 -0
- package/dist/client/utils.js.map +1 -0
- package/dist/component/_generated/api.d.ts +48 -0
- package/dist/component/_generated/api.d.ts.map +1 -0
- package/dist/component/_generated/api.js +31 -0
- package/dist/component/_generated/api.js.map +1 -0
- package/dist/component/_generated/component.d.ts +3120 -0
- package/dist/component/_generated/component.d.ts.map +1 -0
- package/dist/component/_generated/component.js +11 -0
- package/dist/component/_generated/component.js.map +1 -0
- package/dist/component/_generated/dataModel.d.ts +46 -0
- package/dist/component/_generated/dataModel.d.ts.map +1 -0
- package/dist/component/_generated/dataModel.js +11 -0
- package/dist/component/_generated/dataModel.js.map +1 -0
- package/dist/component/_generated/server.d.ts +121 -0
- package/dist/component/_generated/server.d.ts.map +1 -0
- package/dist/component/_generated/server.js +78 -0
- package/dist/component/_generated/server.js.map +1 -0
- package/dist/component/apiKeys.d.ts +11 -0
- package/dist/component/apiKeys.d.ts.map +1 -0
- package/dist/component/apiKeys.js +69 -0
- package/dist/component/apiKeys.js.map +1 -0
- package/dist/component/convex.config.d.ts +3 -0
- package/dist/component/convex.config.d.ts.map +1 -0
- package/dist/component/convex.config.js +3 -0
- package/dist/component/convex.config.js.map +1 -0
- package/dist/component/files.d.ts +97 -0
- package/dist/component/files.d.ts.map +1 -0
- package/dist/component/files.js +190 -0
- package/dist/component/files.js.map +1 -0
- package/dist/component/messages.d.ts +3851 -0
- package/dist/component/messages.d.ts.map +1 -0
- package/dist/component/messages.js +757 -0
- package/dist/component/messages.js.map +1 -0
- package/dist/component/schema.d.ts +8029 -0
- package/dist/component/schema.d.ts.map +1 -0
- package/dist/component/schema.js +147 -0
- package/dist/component/schema.js.map +1 -0
- package/dist/component/streams.d.ts +128 -0
- package/dist/component/streams.d.ts.map +1 -0
- package/dist/component/streams.js +413 -0
- package/dist/component/streams.js.map +1 -0
- package/dist/component/threads.d.ts +115 -0
- package/dist/component/threads.d.ts.map +1 -0
- package/dist/component/threads.js +208 -0
- package/dist/component/threads.js.map +1 -0
- package/dist/component/users.d.ts +52 -0
- package/dist/component/users.d.ts.map +1 -0
- package/dist/component/users.js +229 -0
- package/dist/component/users.js.map +1 -0
- package/dist/component/vector/index.d.ts +61 -0
- package/dist/component/vector/index.d.ts.map +1 -0
- package/dist/component/vector/index.js +146 -0
- package/dist/component/vector/index.js.map +1 -0
- package/dist/component/vector/tables.d.ts +58 -0
- package/dist/component/vector/tables.d.ts.map +1 -0
- package/dist/component/vector/tables.js +56 -0
- package/dist/component/vector/tables.js.map +1 -0
- package/dist/deltas.d.ts +43 -0
- package/dist/deltas.d.ts.map +1 -0
- package/dist/deltas.js +446 -0
- package/dist/deltas.js.map +1 -0
- package/dist/mapping.d.ts +76 -0
- package/dist/mapping.d.ts.map +1 -0
- package/dist/mapping.js +720 -0
- package/dist/mapping.js.map +1 -0
- package/dist/react/SmoothText.d.ts +5 -0
- package/dist/react/SmoothText.d.ts.map +1 -0
- package/dist/react/SmoothText.js +6 -0
- package/dist/react/SmoothText.js.map +1 -0
- package/dist/react/index.d.ts +25 -0
- package/dist/react/index.d.ts.map +1 -0
- package/dist/react/index.js +70 -0
- package/dist/react/index.js.map +1 -0
- package/dist/react/optimisticallySendMessage.d.ts +42 -0
- package/dist/react/optimisticallySendMessage.d.ts.map +1 -0
- package/dist/react/optimisticallySendMessage.js +74 -0
- package/dist/react/optimisticallySendMessage.js.map +1 -0
- package/dist/react/types.d.ts +12 -0
- package/dist/react/types.d.ts.map +1 -0
- package/dist/react/types.js +2 -0
- package/dist/react/types.js.map +1 -0
- package/dist/react/useDeltaStreams.d.ts +10 -0
- package/dist/react/useDeltaStreams.d.ts.map +1 -0
- package/dist/react/useDeltaStreams.js +106 -0
- package/dist/react/useDeltaStreams.js.map +1 -0
- package/dist/react/useSmoothText.d.ts +27 -0
- package/dist/react/useSmoothText.d.ts.map +1 -0
- package/dist/react/useSmoothText.js +68 -0
- package/dist/react/useSmoothText.js.map +1 -0
- package/dist/react/useStreamingUIMessages.d.ts +22 -0
- package/dist/react/useStreamingUIMessages.d.ts.map +1 -0
- package/dist/react/useStreamingUIMessages.js +92 -0
- package/dist/react/useStreamingUIMessages.js.map +1 -0
- package/dist/react/useThreadMessages.d.ts +104 -0
- package/dist/react/useThreadMessages.d.ts.map +1 -0
- package/dist/react/useThreadMessages.js +148 -0
- package/dist/react/useThreadMessages.js.map +1 -0
- package/dist/react/useUIMessages.d.ts +96 -0
- package/dist/react/useUIMessages.d.ts.map +1 -0
- package/dist/react/useUIMessages.js +108 -0
- package/dist/react/useUIMessages.js.map +1 -0
- package/dist/shared.d.ts +26 -0
- package/dist/shared.d.ts.map +1 -0
- package/dist/shared.js +67 -0
- package/dist/shared.js.map +1 -0
- package/dist/validators.d.ts +24516 -0
- package/dist/validators.d.ts.map +1 -0
- package/dist/validators.js +475 -0
- package/dist/validators.js.map +1 -0
- package/package.json +39 -39
- package/src/UIMessages.ts +0 -9
- package/src/client/approval.test.ts +144 -0
- package/src/client/index.ts +17 -9
- package/src/client/search.ts +0 -29
- package/src/client/start.ts +2 -2
- package/src/client/streamText.ts +9 -0
- package/src/client/streaming.ts +30 -23
- package/src/component/_generated/dataModel.ts +1 -1
- package/src/component/_generated/server.ts +0 -5
- package/src/mapping.ts +38 -22
- package/src/node_modules/.vite/vitest/da39a3ee5e6b4b0d3255bfef95601890afd80709/results.json +1 -0
|
@@ -0,0 +1,757 @@
|
|
|
1
|
+
import { assert, omit, pick } from "convex-helpers";
|
|
2
|
+
import { mergedStream, stream } from "convex-helpers/server/stream";
|
|
3
|
+
import { paginationOptsValidator, } from "convex/server";
|
|
4
|
+
import { DEFAULT_MESSAGE_RANGE, DEFAULT_RECENT_MESSAGES, extractText, isTool, sorted, } from "../shared.js";
|
|
5
|
+
import { vMessageDoc, vMessageEmbeddingsWithDimension, vMessageStatus, vMessageWithMetadataInternal, vPaginationResult, } from "../validators.js";
|
|
6
|
+
import { api, internal } from "./_generated/api.js";
|
|
7
|
+
import { action, internalMutation, internalQuery, mutation, query, } from "./_generated/server.js";
|
|
8
|
+
import { schema, v } from "./schema.js";
|
|
9
|
+
import { insertVector, searchVectors } from "./vector/index.js";
|
|
10
|
+
import { validateVectorDimension, vVectorId, } from "./vector/tables.js";
|
|
11
|
+
import { changeRefcount } from "./files.js";
|
|
12
|
+
import { getStreamingMessagesWithMetadata, finishHandler } from "./streams.js";
|
|
13
|
+
import { partial } from "convex-helpers/validators";
|
|
14
|
+
function publicMessage(message) {
|
|
15
|
+
return omit(message, ["parentMessageId", "stepId", "files"]);
|
|
16
|
+
}
|
|
17
|
+
export async function deleteMessage(ctx, messageDoc) {
|
|
18
|
+
await ctx.db.delete(messageDoc._id);
|
|
19
|
+
if (messageDoc.embeddingId) {
|
|
20
|
+
await ctx.db.delete(messageDoc.embeddingId);
|
|
21
|
+
}
|
|
22
|
+
if (messageDoc.fileIds) {
|
|
23
|
+
await changeRefcount(ctx, messageDoc.fileIds, []);
|
|
24
|
+
}
|
|
25
|
+
}
|
|
26
|
+
export const deleteByIds = mutation({
|
|
27
|
+
args: { messageIds: v.array(v.id("messages")) },
|
|
28
|
+
returns: v.array(v.id("messages")),
|
|
29
|
+
handler: async (ctx, args) => {
|
|
30
|
+
const deletedMessageIds = await Promise.all(args.messageIds.map(async (id) => {
|
|
31
|
+
const message = await ctx.db.get(id);
|
|
32
|
+
if (message) {
|
|
33
|
+
await deleteMessage(ctx, message);
|
|
34
|
+
return id;
|
|
35
|
+
}
|
|
36
|
+
return null;
|
|
37
|
+
}));
|
|
38
|
+
return deletedMessageIds.filter((id) => id !== null);
|
|
39
|
+
},
|
|
40
|
+
});
|
|
41
|
+
export const messageStatuses = vMessageDoc.fields.status.members.map((m) => m.value);
|
|
42
|
+
export const deleteByOrder = mutation({
|
|
43
|
+
args: {
|
|
44
|
+
threadId: v.id("threads"),
|
|
45
|
+
startOrder: v.number(),
|
|
46
|
+
startStepOrder: v.optional(v.number()),
|
|
47
|
+
endOrder: v.number(),
|
|
48
|
+
endStepOrder: v.optional(v.number()),
|
|
49
|
+
},
|
|
50
|
+
returns: v.object({
|
|
51
|
+
isDone: v.boolean(),
|
|
52
|
+
lastOrder: v.optional(v.number()),
|
|
53
|
+
lastStepOrder: v.optional(v.number()),
|
|
54
|
+
}),
|
|
55
|
+
handler: async (ctx, args) => {
|
|
56
|
+
const messages = await orderedMessagesStream(ctx, {
|
|
57
|
+
threadId: args.threadId,
|
|
58
|
+
sortOrder: "asc",
|
|
59
|
+
startOrder: args.startOrder,
|
|
60
|
+
startOrderBound: "gte",
|
|
61
|
+
})
|
|
62
|
+
.narrow({
|
|
63
|
+
lowerBound: args.startStepOrder
|
|
64
|
+
? [args.startOrder, args.startStepOrder]
|
|
65
|
+
: [args.startOrder],
|
|
66
|
+
lowerBoundInclusive: true,
|
|
67
|
+
upperBound: args.endStepOrder
|
|
68
|
+
? [args.endOrder, args.endStepOrder]
|
|
69
|
+
: [args.endOrder],
|
|
70
|
+
upperBoundInclusive: false,
|
|
71
|
+
})
|
|
72
|
+
.take(64);
|
|
73
|
+
await Promise.all(messages.map((m) => deleteMessage(ctx, m)));
|
|
74
|
+
return {
|
|
75
|
+
isDone: messages.length < 64,
|
|
76
|
+
lastOrder: messages.at(-1)?.order,
|
|
77
|
+
lastStepOrder: messages.at(-1)?.stepOrder,
|
|
78
|
+
};
|
|
79
|
+
},
|
|
80
|
+
});
|
|
81
|
+
const addMessagesArgs = {
|
|
82
|
+
userId: v.optional(v.string()),
|
|
83
|
+
threadId: v.id("threads"),
|
|
84
|
+
promptMessageId: v.optional(v.id("messages")),
|
|
85
|
+
agentName: v.optional(v.string()),
|
|
86
|
+
messages: v.array(vMessageWithMetadataInternal),
|
|
87
|
+
embeddings: v.optional(vMessageEmbeddingsWithDimension),
|
|
88
|
+
failPendingSteps: v.optional(v.boolean()),
|
|
89
|
+
// A pending message to update. If the pending message failed, abort.
|
|
90
|
+
pendingMessageId: v.optional(v.id("messages")),
|
|
91
|
+
// if set to true, these messages will not show up in text or vector search
|
|
92
|
+
// results for the userId
|
|
93
|
+
hideFromUserIdSearch: v.optional(v.boolean()),
|
|
94
|
+
// If provided, finish this stream atomically with the message save.
|
|
95
|
+
// This prevents UI flickering from separate mutations (issue #181).
|
|
96
|
+
finishStreamId: v.optional(v.id("streamingMessages")),
|
|
97
|
+
};
|
|
98
|
+
export const addMessages = mutation({
|
|
99
|
+
args: addMessagesArgs,
|
|
100
|
+
handler: addMessagesHandler,
|
|
101
|
+
returns: v.object({ messages: v.array(vMessageDoc) }),
|
|
102
|
+
});
|
|
103
|
+
async function addMessagesHandler(ctx, args) {
|
|
104
|
+
let userId = args.userId;
|
|
105
|
+
const threadId = args.threadId;
|
|
106
|
+
if (!userId && args.threadId) {
|
|
107
|
+
const thread = await ctx.db.get(args.threadId);
|
|
108
|
+
assert(thread, `Thread ${args.threadId} not found`);
|
|
109
|
+
userId = thread.userId;
|
|
110
|
+
}
|
|
111
|
+
const { embeddings, failPendingSteps,
|
|
112
|
+
// Destructured separately to exclude from `...rest` (used in addMessages args, not message fields)
|
|
113
|
+
finishStreamId, messages, promptMessageId, pendingMessageId, hideFromUserIdSearch, ...rest } = args;
|
|
114
|
+
const promptMessage = promptMessageId && (await ctx.db.get(promptMessageId));
|
|
115
|
+
if (failPendingSteps) {
|
|
116
|
+
assert(args.threadId, "threadId is required to fail pending steps");
|
|
117
|
+
const pendingMessages = await ctx.db
|
|
118
|
+
.query("messages")
|
|
119
|
+
.withIndex("threadId_status_tool_order_stepOrder", (q) => q.eq("threadId", threadId).eq("status", "pending"))
|
|
120
|
+
.order("desc")
|
|
121
|
+
.take(100);
|
|
122
|
+
await Promise.all(pendingMessages
|
|
123
|
+
.filter((m) => !promptMessage || m.order === promptMessage.order)
|
|
124
|
+
.filter((m) => !pendingMessageId || m._id !== pendingMessageId)
|
|
125
|
+
.map(async (m) => {
|
|
126
|
+
if (m.embeddingId) {
|
|
127
|
+
await ctx.db.delete(m.embeddingId);
|
|
128
|
+
}
|
|
129
|
+
await ctx.db.patch(m._id, {
|
|
130
|
+
status: "failed",
|
|
131
|
+
error: "Restarting",
|
|
132
|
+
embeddingId: undefined,
|
|
133
|
+
});
|
|
134
|
+
}));
|
|
135
|
+
}
|
|
136
|
+
let order, stepOrder;
|
|
137
|
+
let fail = false;
|
|
138
|
+
let error;
|
|
139
|
+
if (promptMessageId) {
|
|
140
|
+
assert(promptMessage, `Parent message ${promptMessageId} not found`);
|
|
141
|
+
if (promptMessage.status === "failed") {
|
|
142
|
+
fail = true;
|
|
143
|
+
error = promptMessage.error ?? error ?? "The prompt message failed";
|
|
144
|
+
}
|
|
145
|
+
order = promptMessage.order;
|
|
146
|
+
// Defend against there being existing messages with this parent.
|
|
147
|
+
const maxMessage = await getMaxMessage(ctx, threadId, order);
|
|
148
|
+
stepOrder = maxMessage?.stepOrder ?? promptMessage.stepOrder;
|
|
149
|
+
}
|
|
150
|
+
else {
|
|
151
|
+
const maxMessage = await getMaxMessage(ctx, threadId);
|
|
152
|
+
order = maxMessage?.order ?? -1;
|
|
153
|
+
stepOrder = maxMessage?.stepOrder ?? -1;
|
|
154
|
+
}
|
|
155
|
+
const toReturn = [];
|
|
156
|
+
if (embeddings) {
|
|
157
|
+
assert(embeddings.vectors.length === messages.length, "embeddings.vectors.length must match messages.length");
|
|
158
|
+
}
|
|
159
|
+
for (let i = 0; i < messages.length; i++) {
|
|
160
|
+
const message = messages[i];
|
|
161
|
+
let embeddingId;
|
|
162
|
+
if (embeddings &&
|
|
163
|
+
embeddings.vectors[i] &&
|
|
164
|
+
!fail &&
|
|
165
|
+
message.status !== "failed") {
|
|
166
|
+
embeddingId = await insertVector(ctx, embeddings.dimension, {
|
|
167
|
+
vector: embeddings.vectors[i],
|
|
168
|
+
model: embeddings.model,
|
|
169
|
+
table: "messages",
|
|
170
|
+
userId: hideFromUserIdSearch ? undefined : userId,
|
|
171
|
+
threadId,
|
|
172
|
+
});
|
|
173
|
+
}
|
|
174
|
+
const messageDoc = {
|
|
175
|
+
...rest,
|
|
176
|
+
...message,
|
|
177
|
+
embeddingId,
|
|
178
|
+
parentMessageId: promptMessageId,
|
|
179
|
+
userId,
|
|
180
|
+
tool: isTool(message.message),
|
|
181
|
+
text: hideFromUserIdSearch ? undefined : extractText(message.message),
|
|
182
|
+
status: fail ? "failed" : (message.status ?? "success"),
|
|
183
|
+
error: fail ? error : message.error,
|
|
184
|
+
};
|
|
185
|
+
// If there is a pending message, we replace that one with the first message
|
|
186
|
+
// and subsequent ones will follow the regular order/subOrder advancement.
|
|
187
|
+
if (i === 0 && pendingMessageId) {
|
|
188
|
+
const pendingMessage = await ctx.db.get(pendingMessageId);
|
|
189
|
+
assert(pendingMessage, `Pending msg ${pendingMessageId} not found`);
|
|
190
|
+
if (pendingMessage.status === "failed") {
|
|
191
|
+
fail = true;
|
|
192
|
+
error =
|
|
193
|
+
`Trying to update a message that failed: ${pendingMessageId}, ` +
|
|
194
|
+
`error: ${pendingMessage.error ?? error}`;
|
|
195
|
+
messageDoc.status = "failed";
|
|
196
|
+
messageDoc.error = error;
|
|
197
|
+
}
|
|
198
|
+
if (message.fileIds) {
|
|
199
|
+
await changeRefcount(ctx, pendingMessage.fileIds ?? [], message.fileIds);
|
|
200
|
+
}
|
|
201
|
+
await ctx.db.replace(pendingMessage._id, {
|
|
202
|
+
...messageDoc,
|
|
203
|
+
order: pendingMessage.order,
|
|
204
|
+
stepOrder: pendingMessage.stepOrder,
|
|
205
|
+
});
|
|
206
|
+
toReturn.push((await ctx.db.get(pendingMessage._id)));
|
|
207
|
+
continue;
|
|
208
|
+
}
|
|
209
|
+
if (message.message.role === "user") {
|
|
210
|
+
if (promptMessage && promptMessage.order === order) {
|
|
211
|
+
// see if there's a later message than the parent message order
|
|
212
|
+
const maxMessage = await getMaxMessage(ctx, threadId);
|
|
213
|
+
order = (maxMessage?.order ?? order) + 1;
|
|
214
|
+
}
|
|
215
|
+
else {
|
|
216
|
+
order++;
|
|
217
|
+
}
|
|
218
|
+
stepOrder = 0;
|
|
219
|
+
}
|
|
220
|
+
else {
|
|
221
|
+
if (order < 0) {
|
|
222
|
+
order = 0;
|
|
223
|
+
}
|
|
224
|
+
stepOrder++;
|
|
225
|
+
}
|
|
226
|
+
const messageId = await ctx.db.insert("messages", {
|
|
227
|
+
...messageDoc,
|
|
228
|
+
order,
|
|
229
|
+
stepOrder,
|
|
230
|
+
});
|
|
231
|
+
if (message.fileIds) {
|
|
232
|
+
await changeRefcount(ctx, [], message.fileIds);
|
|
233
|
+
}
|
|
234
|
+
// TODO: delete the associated stream data for the order/stepOrder
|
|
235
|
+
toReturn.push((await ctx.db.get(messageId)));
|
|
236
|
+
}
|
|
237
|
+
// Atomically finish the stream if requested, preventing UI flickering
|
|
238
|
+
// from separate mutations for message save and stream finish (issue #181).
|
|
239
|
+
if (finishStreamId) {
|
|
240
|
+
await finishHandler(ctx, { streamId: finishStreamId });
|
|
241
|
+
}
|
|
242
|
+
return { messages: toReturn.map(publicMessage) };
|
|
243
|
+
}
|
|
244
|
+
// exported for tests
|
|
245
|
+
export async function getMaxMessage(ctx, threadId, order) {
|
|
246
|
+
return orderedMessagesStream(ctx, {
|
|
247
|
+
threadId,
|
|
248
|
+
sortOrder: "desc",
|
|
249
|
+
startOrder: order,
|
|
250
|
+
startOrderBound: "eq",
|
|
251
|
+
}).first();
|
|
252
|
+
}
|
|
253
|
+
function orderedMessagesStream(ctx, args) {
|
|
254
|
+
return mergedStream([true, false].flatMap((tool) => messageStatuses.map((status) => stream(ctx.db, schema)
|
|
255
|
+
.query("messages")
|
|
256
|
+
.withIndex("threadId_status_tool_order_stepOrder", (q) => {
|
|
257
|
+
const qq = q
|
|
258
|
+
.eq("threadId", args.threadId)
|
|
259
|
+
.eq("status", status)
|
|
260
|
+
.eq("tool", tool);
|
|
261
|
+
if (args.startOrder !== undefined) {
|
|
262
|
+
if (args.startOrderBound === "gte") {
|
|
263
|
+
return qq.gte("order", args.startOrder);
|
|
264
|
+
}
|
|
265
|
+
else {
|
|
266
|
+
return qq.eq("order", args.startOrder);
|
|
267
|
+
}
|
|
268
|
+
}
|
|
269
|
+
return qq;
|
|
270
|
+
})
|
|
271
|
+
.order(args.sortOrder))), ["order", "stepOrder"]);
|
|
272
|
+
}
|
|
273
|
+
export const finalizeMessage = mutation({
|
|
274
|
+
args: {
|
|
275
|
+
messageId: v.id("messages"),
|
|
276
|
+
result: v.union(v.object({ status: v.literal("success") }), v.object({ status: v.literal("failed"), error: v.string() })),
|
|
277
|
+
},
|
|
278
|
+
returns: v.null(),
|
|
279
|
+
handler: async (ctx, { messageId, result }) => {
|
|
280
|
+
const message = await ctx.db.get(messageId);
|
|
281
|
+
assert(message, `Message ${messageId} not found`);
|
|
282
|
+
if (message.status !== "pending") {
|
|
283
|
+
console.debug("Trying to finalize a message that's already", message.status);
|
|
284
|
+
return;
|
|
285
|
+
}
|
|
286
|
+
// See if we can add any in-progress data
|
|
287
|
+
if (!message.message?.content.length) {
|
|
288
|
+
const messages = await getStreamingMessagesWithMetadata(ctx, message, result);
|
|
289
|
+
if (messages.length > 0) {
|
|
290
|
+
await addMessagesHandler(ctx, {
|
|
291
|
+
messages,
|
|
292
|
+
threadId: message.threadId,
|
|
293
|
+
agentName: message.agentName,
|
|
294
|
+
failPendingSteps: false,
|
|
295
|
+
pendingMessageId: messageId,
|
|
296
|
+
userId: message.userId,
|
|
297
|
+
embeddings: undefined,
|
|
298
|
+
});
|
|
299
|
+
return;
|
|
300
|
+
}
|
|
301
|
+
}
|
|
302
|
+
if (result.status === "failed") {
|
|
303
|
+
if (message.embeddingId) {
|
|
304
|
+
await ctx.db.delete(message.embeddingId);
|
|
305
|
+
}
|
|
306
|
+
await ctx.db.patch(messageId, {
|
|
307
|
+
status: "failed",
|
|
308
|
+
error: result.error,
|
|
309
|
+
embeddingId: undefined,
|
|
310
|
+
});
|
|
311
|
+
}
|
|
312
|
+
else {
|
|
313
|
+
await ctx.db.patch(messageId, { status: "success" });
|
|
314
|
+
}
|
|
315
|
+
},
|
|
316
|
+
});
|
|
317
|
+
export const updateMessage = mutation({
|
|
318
|
+
args: {
|
|
319
|
+
messageId: v.id("messages"),
|
|
320
|
+
patch: v.object(partial(pick(schema.tables.messages.validator.fields, [
|
|
321
|
+
"message",
|
|
322
|
+
"fileIds",
|
|
323
|
+
"status",
|
|
324
|
+
"error",
|
|
325
|
+
"model",
|
|
326
|
+
"provider",
|
|
327
|
+
"providerOptions",
|
|
328
|
+
"finishReason",
|
|
329
|
+
]))),
|
|
330
|
+
},
|
|
331
|
+
returns: vMessageDoc,
|
|
332
|
+
handler: async (ctx, args) => {
|
|
333
|
+
const message = await ctx.db.get(args.messageId);
|
|
334
|
+
assert(message, `Message ${args.messageId} not found`);
|
|
335
|
+
if (args.patch.fileIds) {
|
|
336
|
+
await changeRefcount(ctx, message.fileIds ?? [], args.patch.fileIds);
|
|
337
|
+
}
|
|
338
|
+
const patch = { ...args.patch };
|
|
339
|
+
if (args.patch.message !== undefined) {
|
|
340
|
+
patch.message = args.patch.message;
|
|
341
|
+
patch.tool = isTool(args.patch.message);
|
|
342
|
+
patch.text = extractText(args.patch.message);
|
|
343
|
+
}
|
|
344
|
+
if (args.patch.status === "failed") {
|
|
345
|
+
if (message.embeddingId) {
|
|
346
|
+
await ctx.db.delete(message.embeddingId);
|
|
347
|
+
}
|
|
348
|
+
patch.embeddingId = undefined;
|
|
349
|
+
}
|
|
350
|
+
await ctx.db.patch(args.messageId, patch);
|
|
351
|
+
return publicMessage((await ctx.db.get(args.messageId)));
|
|
352
|
+
},
|
|
353
|
+
});
|
|
354
|
+
const cloneMessageArgs = {
|
|
355
|
+
sourceThreadId: v.id("threads"),
|
|
356
|
+
targetThreadId: v.id("threads"),
|
|
357
|
+
// defaults to false, so searching for a message by userId will not find
|
|
358
|
+
// these copies
|
|
359
|
+
copyUserIdForVectorSearch: v.optional(v.boolean()),
|
|
360
|
+
// defaults to false, so tool calls & responses will be copied
|
|
361
|
+
excludeToolMessages: v.optional(v.boolean()),
|
|
362
|
+
// defaults to copying all messages, but you could just copy success messages.
|
|
363
|
+
statuses: v.optional(v.array(vMessageStatus)),
|
|
364
|
+
// stop at this message id
|
|
365
|
+
upToAndIncludingMessageId: v.optional(v.id("messages")),
|
|
366
|
+
// defaults to 0. the messages will be inserted starting at this order.
|
|
367
|
+
insertAtOrder: v.optional(v.number()),
|
|
368
|
+
};
|
|
369
|
+
export const cloneMessageBatch = internalMutation({
|
|
370
|
+
args: {
|
|
371
|
+
...cloneMessageArgs,
|
|
372
|
+
paginationOpts: paginationOptsValidator,
|
|
373
|
+
},
|
|
374
|
+
handler: async (ctx, args) => {
|
|
375
|
+
const orderOffset = args.insertAtOrder ?? 0;
|
|
376
|
+
const result = await listMessagesByThreadIdHandler(ctx, {
|
|
377
|
+
threadId: args.sourceThreadId,
|
|
378
|
+
excludeToolMessages: args.excludeToolMessages,
|
|
379
|
+
order: "desc",
|
|
380
|
+
paginationOpts: args.paginationOpts,
|
|
381
|
+
statuses: args.statuses,
|
|
382
|
+
upToAndIncludingMessageId: args.upToAndIncludingMessageId,
|
|
383
|
+
});
|
|
384
|
+
const existing = result.page.length === 0
|
|
385
|
+
? []
|
|
386
|
+
: await mergedStream([true, false].flatMap((tool) => messageStatuses.map((status) => stream(ctx.db, schema)
|
|
387
|
+
.query("messages")
|
|
388
|
+
.withIndex("threadId_status_tool_order_stepOrder", (q) => q
|
|
389
|
+
.eq("threadId", args.targetThreadId)
|
|
390
|
+
.eq("status", status)
|
|
391
|
+
.eq("tool", tool)
|
|
392
|
+
.gte("order", result.page[0].order)
|
|
393
|
+
.lte("order", result.page[result.page.length - 1].order)))), ["order", "stepOrder"]).collect();
|
|
394
|
+
await Promise.all(result.page
|
|
395
|
+
.filter((m) => !existing.some((e) => e.order === m.order && e.stepOrder === m.stepOrder))
|
|
396
|
+
.map(async (m) => {
|
|
397
|
+
// update file refs
|
|
398
|
+
if (m.fileIds) {
|
|
399
|
+
await changeRefcount(ctx, [], m.fileIds);
|
|
400
|
+
}
|
|
401
|
+
let embeddingId = undefined;
|
|
402
|
+
if (m.embeddingId) {
|
|
403
|
+
const vector = await ctx.db.get(m.embeddingId);
|
|
404
|
+
assert(vector, `Vector ${m.embeddingId} not found`);
|
|
405
|
+
const dimension = vector.vector.length;
|
|
406
|
+
validateVectorDimension(dimension);
|
|
407
|
+
embeddingId = await insertVector(ctx, dimension, {
|
|
408
|
+
...pick(vector, ["model", "table", "vector"]),
|
|
409
|
+
userId: args.copyUserIdForVectorSearch
|
|
410
|
+
? vector.userId
|
|
411
|
+
: undefined,
|
|
412
|
+
threadId: args.targetThreadId,
|
|
413
|
+
});
|
|
414
|
+
}
|
|
415
|
+
await ctx.db.insert("messages", {
|
|
416
|
+
...omit(m, [
|
|
417
|
+
"_id",
|
|
418
|
+
"_creationTime",
|
|
419
|
+
"threadId",
|
|
420
|
+
"order",
|
|
421
|
+
"embeddingId",
|
|
422
|
+
]),
|
|
423
|
+
embeddingId,
|
|
424
|
+
threadId: args.targetThreadId,
|
|
425
|
+
order: orderOffset + m.order,
|
|
426
|
+
});
|
|
427
|
+
}));
|
|
428
|
+
return {
|
|
429
|
+
numCopied: result.page.length,
|
|
430
|
+
continueCursor: result.continueCursor,
|
|
431
|
+
isDone: result.isDone,
|
|
432
|
+
};
|
|
433
|
+
},
|
|
434
|
+
});
|
|
435
|
+
export const cloneThread = action({
|
|
436
|
+
args: {
|
|
437
|
+
...cloneMessageArgs,
|
|
438
|
+
batchSize: v.optional(v.number()),
|
|
439
|
+
// how many messages to copy
|
|
440
|
+
limit: v.optional(v.number()),
|
|
441
|
+
},
|
|
442
|
+
returns: v.number(),
|
|
443
|
+
handler: async (ctx, args) => {
|
|
444
|
+
let cursor = null;
|
|
445
|
+
let copiedSoFar = 0;
|
|
446
|
+
while (copiedSoFar < (args.limit ?? Infinity)) {
|
|
447
|
+
const numToCopy = Math.min(args.batchSize ?? DEFAULT_RECENT_MESSAGES, args.limit ?? Infinity - copiedSoFar);
|
|
448
|
+
const result = await ctx.runMutation(internal.messages.cloneMessageBatch, {
|
|
449
|
+
...args,
|
|
450
|
+
paginationOpts: {
|
|
451
|
+
cursor,
|
|
452
|
+
numItems: numToCopy,
|
|
453
|
+
},
|
|
454
|
+
});
|
|
455
|
+
copiedSoFar += result.numCopied;
|
|
456
|
+
cursor = result.continueCursor;
|
|
457
|
+
if (result.isDone) {
|
|
458
|
+
break;
|
|
459
|
+
}
|
|
460
|
+
}
|
|
461
|
+
return copiedSoFar;
|
|
462
|
+
},
|
|
463
|
+
});
|
|
464
|
+
export const listMessagesByThreadIdArgs = {
|
|
465
|
+
threadId: v.id("threads"),
|
|
466
|
+
excludeToolMessages: v.optional(v.boolean()),
|
|
467
|
+
/** What order to sort the messages in. To get the latest, use "desc". */
|
|
468
|
+
order: v.union(v.literal("asc"), v.literal("desc")),
|
|
469
|
+
paginationOpts: v.optional(paginationOptsValidator),
|
|
470
|
+
statuses: v.optional(v.array(vMessageStatus)),
|
|
471
|
+
upToAndIncludingMessageId: v.optional(v.id("messages")),
|
|
472
|
+
};
|
|
473
|
+
export const listMessagesByThreadId = query({
|
|
474
|
+
args: listMessagesByThreadIdArgs,
|
|
475
|
+
handler: async (ctx, args) => {
|
|
476
|
+
const messages = await listMessagesByThreadIdHandler(ctx, args);
|
|
477
|
+
return { ...messages, page: messages.page.map(publicMessage) };
|
|
478
|
+
},
|
|
479
|
+
returns: vPaginationResult(vMessageDoc),
|
|
480
|
+
});
|
|
481
|
+
async function listMessagesByThreadIdHandler(ctx, args) {
|
|
482
|
+
const statuses = args.statuses ?? vMessageStatus.members.map((m) => m.value);
|
|
483
|
+
const last = args.upToAndIncludingMessageId &&
|
|
484
|
+
(await ctx.db.get(args.upToAndIncludingMessageId));
|
|
485
|
+
assert(!last || last.threadId === args.threadId, "upToAndIncludingMessageId must be a message in the thread");
|
|
486
|
+
const toolOptions = args.excludeToolMessages ? [false] : [true, false];
|
|
487
|
+
const order = args.order ?? "desc";
|
|
488
|
+
const streams = toolOptions.flatMap((tool) => statuses.map((status) => stream(ctx.db, schema)
|
|
489
|
+
.query("messages")
|
|
490
|
+
.withIndex("threadId_status_tool_order_stepOrder", (q) => {
|
|
491
|
+
const qq = q
|
|
492
|
+
.eq("threadId", args.threadId)
|
|
493
|
+
.eq("status", status)
|
|
494
|
+
.eq("tool", tool);
|
|
495
|
+
if (last) {
|
|
496
|
+
return qq.lte("order", last.order);
|
|
497
|
+
}
|
|
498
|
+
return qq;
|
|
499
|
+
})
|
|
500
|
+
.order(order)
|
|
501
|
+
.filterWith(
|
|
502
|
+
// We allow all messages on the same order.
|
|
503
|
+
async (m) => !last || m.order <= last.order)));
|
|
504
|
+
const messages = await mergedStream(streams, ["order", "stepOrder"]).paginate(args.paginationOpts ?? {
|
|
505
|
+
numItems: DEFAULT_RECENT_MESSAGES,
|
|
506
|
+
cursor: null,
|
|
507
|
+
});
|
|
508
|
+
if (messages.page.length === 0) {
|
|
509
|
+
messages.isDone = true;
|
|
510
|
+
}
|
|
511
|
+
return messages;
|
|
512
|
+
}
|
|
513
|
+
export const getMessagesByIds = query({
|
|
514
|
+
args: { messageIds: v.array(v.id("messages")) },
|
|
515
|
+
handler: async (ctx, args) => {
|
|
516
|
+
return (await Promise.all(args.messageIds.map((id) => ctx.db.get(id)))).map((m) => (m ? publicMessage(m) : null));
|
|
517
|
+
},
|
|
518
|
+
returns: v.array(v.union(v.null(), vMessageDoc)),
|
|
519
|
+
});
|
|
520
|
+
export const searchMessages = action({
|
|
521
|
+
args: {
|
|
522
|
+
threadId: v.optional(v.id("threads")),
|
|
523
|
+
searchAllMessagesForUserId: v.optional(v.string()),
|
|
524
|
+
targetMessageId: v.optional(v.id("messages")),
|
|
525
|
+
embedding: v.optional(v.array(v.number())),
|
|
526
|
+
embeddingModel: v.optional(v.string()),
|
|
527
|
+
text: v.optional(v.string()),
|
|
528
|
+
textSearch: v.optional(v.boolean()),
|
|
529
|
+
vectorSearch: v.optional(v.boolean()),
|
|
530
|
+
limit: v.number(),
|
|
531
|
+
vectorScoreThreshold: v.optional(v.number()),
|
|
532
|
+
messageRange: v.optional(v.object({ before: v.number(), after: v.number() })),
|
|
533
|
+
},
|
|
534
|
+
returns: v.array(vMessageDoc),
|
|
535
|
+
handler: async (ctx, args) => {
|
|
536
|
+
assert(args.searchAllMessagesForUserId || args.threadId, "Specify userId or threadId");
|
|
537
|
+
const limit = args.limit;
|
|
538
|
+
let textSearchMessages;
|
|
539
|
+
if (args.textSearch) {
|
|
540
|
+
textSearchMessages = await ctx.runQuery(api.messages.textSearch, {
|
|
541
|
+
searchAllMessagesForUserId: args.searchAllMessagesForUserId,
|
|
542
|
+
threadId: args.threadId,
|
|
543
|
+
targetMessageId: args.targetMessageId,
|
|
544
|
+
text: args.text,
|
|
545
|
+
limit,
|
|
546
|
+
});
|
|
547
|
+
}
|
|
548
|
+
if (args.vectorSearch) {
|
|
549
|
+
let embedding = args.embedding;
|
|
550
|
+
let model = args.embeddingModel;
|
|
551
|
+
if (!embedding) {
|
|
552
|
+
if (args.targetMessageId) {
|
|
553
|
+
const target = await ctx.runQuery(api.messages.getMessageSearchFields, {
|
|
554
|
+
messageId: args.targetMessageId,
|
|
555
|
+
});
|
|
556
|
+
assert(target, "Target message embedding not found.");
|
|
557
|
+
embedding = target.embedding;
|
|
558
|
+
model = target.embeddingModel;
|
|
559
|
+
}
|
|
560
|
+
}
|
|
561
|
+
assert(embedding && model, "Embedding missing");
|
|
562
|
+
const dimension = embedding.length;
|
|
563
|
+
validateVectorDimension(dimension);
|
|
564
|
+
const vectors = (await searchVectors(ctx, embedding, {
|
|
565
|
+
dimension,
|
|
566
|
+
model,
|
|
567
|
+
table: "messages",
|
|
568
|
+
searchAllMessagesForUserId: args.searchAllMessagesForUserId,
|
|
569
|
+
threadId: args.threadId,
|
|
570
|
+
limit,
|
|
571
|
+
})).filter((v) => v._score > (args.vectorScoreThreshold ?? 0));
|
|
572
|
+
// Reciprocal rank fusion
|
|
573
|
+
const k = 10;
|
|
574
|
+
const textEmbeddingIds = textSearchMessages?.map((m) => m.embeddingId);
|
|
575
|
+
const vectorScores = vectors
|
|
576
|
+
.map((v, i) => ({
|
|
577
|
+
id: v._id,
|
|
578
|
+
score: 1 / (i + k) +
|
|
579
|
+
1 / ((textEmbeddingIds?.indexOf(v._id) ?? Infinity) + k),
|
|
580
|
+
}))
|
|
581
|
+
.sort((a, b) => b.score - a.score);
|
|
582
|
+
const embeddingIds = vectorScores.slice(0, limit).map((v) => v.id);
|
|
583
|
+
const messages = await ctx.runQuery(internal.messages._fetchSearchMessages, {
|
|
584
|
+
searchAllMessagesForUserId: args.searchAllMessagesForUserId,
|
|
585
|
+
threadId: args.threadId,
|
|
586
|
+
embeddingIds,
|
|
587
|
+
textSearchMessages: textSearchMessages?.filter((m) => !embeddingIds.includes(m.embeddingId)),
|
|
588
|
+
messageRange: args.messageRange ?? DEFAULT_MESSAGE_RANGE,
|
|
589
|
+
beforeMessageId: args.targetMessageId,
|
|
590
|
+
limit,
|
|
591
|
+
});
|
|
592
|
+
return messages;
|
|
593
|
+
}
|
|
594
|
+
return textSearchMessages?.flat() ?? [];
|
|
595
|
+
},
|
|
596
|
+
});
|
|
597
|
+
export const _fetchSearchMessages = internalQuery({
|
|
598
|
+
args: {
|
|
599
|
+
threadId: v.optional(v.id("threads")),
|
|
600
|
+
embeddingIds: v.array(vVectorId),
|
|
601
|
+
searchAllMessagesForUserId: v.optional(v.string()),
|
|
602
|
+
textSearchMessages: v.optional(v.array(vMessageDoc)),
|
|
603
|
+
messageRange: v.object({ before: v.number(), after: v.number() }),
|
|
604
|
+
beforeMessageId: v.optional(v.id("messages")),
|
|
605
|
+
limit: v.number(),
|
|
606
|
+
},
|
|
607
|
+
returns: v.array(vMessageDoc),
|
|
608
|
+
handler: async (ctx, args) => {
|
|
609
|
+
const beforeMessage = args.beforeMessageId && (await ctx.db.get(args.beforeMessageId));
|
|
610
|
+
const { searchAllMessagesForUserId, threadId } = args;
|
|
611
|
+
assert(searchAllMessagesForUserId || threadId, "Specify searchAllMessagesForUserId or threadId to search");
|
|
612
|
+
let messages = (await Promise.all(args.embeddingIds.map((embeddingId) => ctx.db
|
|
613
|
+
.query("messages")
|
|
614
|
+
.withIndex("embeddingId_threadId", (q) => searchAllMessagesForUserId
|
|
615
|
+
? q.eq("embeddingId", embeddingId)
|
|
616
|
+
: q.eq("embeddingId", embeddingId).eq("threadId", threadId))
|
|
617
|
+
.filter((q) => q.and(q.eq(q.field("status"), "success"), searchAllMessagesForUserId
|
|
618
|
+
? q.eq(q.field("userId"), searchAllMessagesForUserId)
|
|
619
|
+
: q.eq(q.field("threadId"), threadId)))
|
|
620
|
+
.first())))
|
|
621
|
+
.filter((m) => m !== undefined &&
|
|
622
|
+
m !== null &&
|
|
623
|
+
!m.tool &&
|
|
624
|
+
(!beforeMessage ||
|
|
625
|
+
m.order < beforeMessage.order ||
|
|
626
|
+
(m.order === beforeMessage.order &&
|
|
627
|
+
m.stepOrder < beforeMessage.stepOrder)))
|
|
628
|
+
.map(publicMessage);
|
|
629
|
+
messages.push(...(args.textSearchMessages ?? []));
|
|
630
|
+
// TODO: prioritize more recent messages
|
|
631
|
+
messages = sorted(messages);
|
|
632
|
+
messages = messages.slice(0, args.limit);
|
|
633
|
+
// Fetch the surrounding messages
|
|
634
|
+
if (!threadId) {
|
|
635
|
+
return messages;
|
|
636
|
+
}
|
|
637
|
+
const included = {};
|
|
638
|
+
for (const m of messages) {
|
|
639
|
+
const searchId = m.threadId ?? m.userId;
|
|
640
|
+
if (!included[searchId]) {
|
|
641
|
+
included[searchId] = new Set();
|
|
642
|
+
}
|
|
643
|
+
included[searchId].add(m.order);
|
|
644
|
+
}
|
|
645
|
+
const ranges = {};
|
|
646
|
+
const { before, after } = args.messageRange;
|
|
647
|
+
for (const m of messages) {
|
|
648
|
+
const searchId = m.threadId ?? m.userId;
|
|
649
|
+
const order = m.order;
|
|
650
|
+
let earliest = order - before;
|
|
651
|
+
let latest = order + after;
|
|
652
|
+
for (; earliest <= latest; earliest++) {
|
|
653
|
+
if (!included[searchId].has(earliest)) {
|
|
654
|
+
break;
|
|
655
|
+
}
|
|
656
|
+
}
|
|
657
|
+
for (; latest >= earliest; latest--) {
|
|
658
|
+
if (!included[searchId].has(latest)) {
|
|
659
|
+
break;
|
|
660
|
+
}
|
|
661
|
+
}
|
|
662
|
+
for (let i = earliest; i <= latest; i++) {
|
|
663
|
+
included[searchId].add(i);
|
|
664
|
+
}
|
|
665
|
+
if (earliest !== latest) {
|
|
666
|
+
const surrounding = await ctx.db
|
|
667
|
+
.query("messages")
|
|
668
|
+
.withIndex("threadId_status_tool_order_stepOrder", (q) => q
|
|
669
|
+
.eq("threadId", m.threadId)
|
|
670
|
+
.eq("status", "success")
|
|
671
|
+
.eq("tool", false)
|
|
672
|
+
.gte("order", earliest)
|
|
673
|
+
.lte("order", latest))
|
|
674
|
+
.collect();
|
|
675
|
+
if (!ranges[searchId]) {
|
|
676
|
+
ranges[searchId] = [];
|
|
677
|
+
}
|
|
678
|
+
ranges[searchId].push(...surrounding);
|
|
679
|
+
}
|
|
680
|
+
}
|
|
681
|
+
for (const r of Object.values(ranges).flat()) {
|
|
682
|
+
if (!messages.some((m) => m._id === r._id)) {
|
|
683
|
+
messages.push(publicMessage(r));
|
|
684
|
+
}
|
|
685
|
+
}
|
|
686
|
+
return sorted(messages);
|
|
687
|
+
},
|
|
688
|
+
});
|
|
689
|
+
// returns ranges of messages in order of text search relevance,
|
|
690
|
+
// excluding duplicates in later ranges.
|
|
691
|
+
export const textSearch = query({
|
|
692
|
+
args: {
|
|
693
|
+
threadId: v.optional(v.id("threads")),
|
|
694
|
+
searchAllMessagesForUserId: v.optional(v.string()),
|
|
695
|
+
text: v.optional(v.string()),
|
|
696
|
+
targetMessageId: v.optional(v.id("messages")),
|
|
697
|
+
limit: v.number(),
|
|
698
|
+
},
|
|
699
|
+
handler: async (ctx, args) => {
|
|
700
|
+
assert(args.searchAllMessagesForUserId || args.threadId, "Specify userId or threadId");
|
|
701
|
+
const targetMessage = args.targetMessageId && (await ctx.db.get(args.targetMessageId));
|
|
702
|
+
const order = targetMessage?.order;
|
|
703
|
+
const text = args.text || targetMessage?.text;
|
|
704
|
+
if (!text) {
|
|
705
|
+
console.warn("No text to search", targetMessage, args.text);
|
|
706
|
+
return [];
|
|
707
|
+
}
|
|
708
|
+
const messages = await ctx.db
|
|
709
|
+
.query("messages")
|
|
710
|
+
.withSearchIndex("text_search", (q) => args.searchAllMessagesForUserId
|
|
711
|
+
? q.search("text", text).eq("userId", args.searchAllMessagesForUserId)
|
|
712
|
+
: q.search("text", text).eq("threadId", args.threadId))
|
|
713
|
+
// Just in case tool messages slip through
|
|
714
|
+
.filter((q) => {
|
|
715
|
+
const qq = q.eq(q.field("tool"), false);
|
|
716
|
+
if (order) {
|
|
717
|
+
return q.and(qq, q.lte(q.field("order"), order));
|
|
718
|
+
}
|
|
719
|
+
return qq;
|
|
720
|
+
})
|
|
721
|
+
.take(args.limit);
|
|
722
|
+
return messages
|
|
723
|
+
.filter((m) => !targetMessage ||
|
|
724
|
+
m.order < targetMessage.order ||
|
|
725
|
+
(m.order === targetMessage.order &&
|
|
726
|
+
m.stepOrder < targetMessage.stepOrder))
|
|
727
|
+
.map(publicMessage);
|
|
728
|
+
},
|
|
729
|
+
returns: v.array(vMessageDoc),
|
|
730
|
+
});
|
|
731
|
+
export const getMessageSearchFields = query({
|
|
732
|
+
args: {
|
|
733
|
+
messageId: v.id("messages"),
|
|
734
|
+
},
|
|
735
|
+
returns: v.object({
|
|
736
|
+
text: v.optional(v.string()),
|
|
737
|
+
embedding: v.optional(v.array(v.number())),
|
|
738
|
+
embeddingModel: v.optional(v.string()),
|
|
739
|
+
}),
|
|
740
|
+
handler: async (ctx, args) => {
|
|
741
|
+
const message = await ctx.db.get(args.messageId);
|
|
742
|
+
const text = message?.text;
|
|
743
|
+
let embedding = undefined;
|
|
744
|
+
let embeddingModel = undefined;
|
|
745
|
+
if (message?.embeddingId) {
|
|
746
|
+
const target = await ctx.db.get(message.embeddingId);
|
|
747
|
+
embedding = target?.vector;
|
|
748
|
+
embeddingModel = target?.model;
|
|
749
|
+
}
|
|
750
|
+
return {
|
|
751
|
+
text,
|
|
752
|
+
embedding,
|
|
753
|
+
embeddingModel,
|
|
754
|
+
};
|
|
755
|
+
},
|
|
756
|
+
});
|
|
757
|
+
//# sourceMappingURL=messages.js.map
|