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