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