@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,15 +1,16 @@
|
|
|
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 { paginationOptsValidator } from "convex/server";
|
|
4
|
-
import { DEFAULT_MESSAGE_RANGE, DEFAULT_RECENT_MESSAGES, extractText, isTool, } from "../shared.js";
|
|
5
|
-
import {
|
|
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
6
|
import { api, internal } from "./_generated/api.js";
|
|
7
|
-
import { action, internalQuery, mutation, query, } from "./_generated/server.js";
|
|
8
|
-
import { schema, v
|
|
9
|
-
import { getThread as _getThread, listThreadsByUserId as _listThreadsByUserId, updateThread as _updateThread, } from "./threads.js";
|
|
7
|
+
import { action, internalMutation, internalQuery, mutation, query, } from "./_generated/server.js";
|
|
8
|
+
import { schema, v } from "./schema.js";
|
|
10
9
|
import { insertVector, searchVectors } from "./vector/index.js";
|
|
11
|
-
import {
|
|
10
|
+
import { validateVectorDimension, vVectorId, } from "./vector/tables.js";
|
|
12
11
|
import { changeRefcount } from "./files.js";
|
|
12
|
+
import { getStreamingMessagesWithMetadata, finishHandler } from "./streams.js";
|
|
13
|
+
import { partial } from "convex-helpers/validators";
|
|
13
14
|
function publicMessage(message) {
|
|
14
15
|
return omit(message, ["parentMessageId", "stepId", "files"]);
|
|
15
16
|
}
|
|
@@ -23,9 +24,7 @@ export async function deleteMessage(ctx, messageDoc) {
|
|
|
23
24
|
}
|
|
24
25
|
}
|
|
25
26
|
export const deleteByIds = mutation({
|
|
26
|
-
args: {
|
|
27
|
-
messageIds: v.array(v.id("messages")),
|
|
28
|
-
},
|
|
27
|
+
args: { messageIds: v.array(v.id("messages")) },
|
|
29
28
|
returns: v.array(v.id("messages")),
|
|
30
29
|
handler: async (ctx, args) => {
|
|
31
30
|
const deletedMessageIds = await Promise.all(args.messageIds.map(async (id) => {
|
|
@@ -54,7 +53,12 @@ export const deleteByOrder = mutation({
|
|
|
54
53
|
lastStepOrder: v.optional(v.number()),
|
|
55
54
|
}),
|
|
56
55
|
handler: async (ctx, args) => {
|
|
57
|
-
const messages = await orderedMessagesStream(ctx,
|
|
56
|
+
const messages = await orderedMessagesStream(ctx, {
|
|
57
|
+
threadId: args.threadId,
|
|
58
|
+
sortOrder: "asc",
|
|
59
|
+
startOrder: args.startOrder,
|
|
60
|
+
startOrderBound: "gte",
|
|
61
|
+
})
|
|
58
62
|
.narrow({
|
|
59
63
|
lowerBound: args.startStepOrder
|
|
60
64
|
? [args.startOrder, args.startStepOrder]
|
|
@@ -80,16 +84,21 @@ const addMessagesArgs = {
|
|
|
80
84
|
promptMessageId: v.optional(v.id("messages")),
|
|
81
85
|
agentName: v.optional(v.string()),
|
|
82
86
|
messages: v.array(vMessageWithMetadataInternal),
|
|
83
|
-
embeddings: v.optional(
|
|
84
|
-
pending: v.optional(v.boolean()),
|
|
87
|
+
embeddings: v.optional(vMessageEmbeddingsWithDimension),
|
|
85
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")),
|
|
86
97
|
};
|
|
87
98
|
export const addMessages = mutation({
|
|
88
99
|
args: addMessagesArgs,
|
|
89
100
|
handler: addMessagesHandler,
|
|
90
|
-
returns: v.object({
|
|
91
|
-
messages: v.array(vMessageDoc),
|
|
92
|
-
}),
|
|
101
|
+
returns: v.object({ messages: v.array(vMessageDoc) }),
|
|
93
102
|
});
|
|
94
103
|
async function addMessagesHandler(ctx, args) {
|
|
95
104
|
let userId = args.userId;
|
|
@@ -99,34 +108,49 @@ async function addMessagesHandler(ctx, args) {
|
|
|
99
108
|
assert(thread, `Thread ${args.threadId} not found`);
|
|
100
109
|
userId = thread.userId;
|
|
101
110
|
}
|
|
102
|
-
const { embeddings, failPendingSteps,
|
|
103
|
-
|
|
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));
|
|
104
115
|
if (failPendingSteps) {
|
|
105
116
|
assert(args.threadId, "threadId is required to fail pending steps");
|
|
106
117
|
const pendingMessages = await ctx.db
|
|
107
118
|
.query("messages")
|
|
108
119
|
.withIndex("threadId_status_tool_order_stepOrder", (q) => q.eq("threadId", threadId).eq("status", "pending"))
|
|
109
|
-
.
|
|
120
|
+
.order("desc")
|
|
121
|
+
.take(100);
|
|
110
122
|
await Promise.all(pendingMessages
|
|
111
|
-
.filter((m) => !
|
|
112
|
-
.
|
|
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
|
+
}));
|
|
113
135
|
}
|
|
114
136
|
let order, stepOrder;
|
|
115
137
|
let fail = false;
|
|
138
|
+
let error;
|
|
116
139
|
if (promptMessageId) {
|
|
117
|
-
assert(
|
|
118
|
-
if (
|
|
140
|
+
assert(promptMessage, `Parent message ${promptMessageId} not found`);
|
|
141
|
+
if (promptMessage.status === "failed") {
|
|
119
142
|
fail = true;
|
|
143
|
+
error = promptMessage.error ?? error ?? "The prompt message failed";
|
|
120
144
|
}
|
|
121
|
-
order =
|
|
145
|
+
order = promptMessage.order;
|
|
122
146
|
// Defend against there being existing messages with this parent.
|
|
123
147
|
const maxMessage = await getMaxMessage(ctx, threadId, order);
|
|
124
|
-
stepOrder = maxMessage?.stepOrder ??
|
|
148
|
+
stepOrder = maxMessage?.stepOrder ?? promptMessage.stepOrder;
|
|
125
149
|
}
|
|
126
150
|
else {
|
|
127
151
|
const maxMessage = await getMaxMessage(ctx, threadId);
|
|
128
|
-
order = maxMessage
|
|
129
|
-
stepOrder = -1;
|
|
152
|
+
order = maxMessage?.order ?? -1;
|
|
153
|
+
stepOrder = maxMessage?.stepOrder ?? -1;
|
|
130
154
|
}
|
|
131
155
|
const toReturn = [];
|
|
132
156
|
if (embeddings) {
|
|
@@ -135,99 +159,174 @@ async function addMessagesHandler(ctx, args) {
|
|
|
135
159
|
for (let i = 0; i < messages.length; i++) {
|
|
136
160
|
const message = messages[i];
|
|
137
161
|
let embeddingId;
|
|
138
|
-
if (embeddings &&
|
|
162
|
+
if (embeddings &&
|
|
163
|
+
embeddings.vectors[i] &&
|
|
164
|
+
!fail &&
|
|
165
|
+
message.status !== "failed") {
|
|
139
166
|
embeddingId = await insertVector(ctx, embeddings.dimension, {
|
|
140
167
|
vector: embeddings.vectors[i],
|
|
141
168
|
model: embeddings.model,
|
|
142
169
|
table: "messages",
|
|
143
|
-
userId,
|
|
170
|
+
userId: hideFromUserIdSearch ? undefined : userId,
|
|
144
171
|
threadId,
|
|
145
172
|
});
|
|
146
173
|
}
|
|
147
|
-
|
|
148
|
-
const messageId = await ctx.db.insert("messages", {
|
|
174
|
+
const messageDoc = {
|
|
149
175
|
...rest,
|
|
150
176
|
...message,
|
|
151
177
|
embeddingId,
|
|
152
178
|
parentMessageId: promptMessageId,
|
|
153
179
|
userId,
|
|
154
|
-
order,
|
|
155
180
|
tool: isTool(message.message),
|
|
156
|
-
text: extractText(message.message),
|
|
157
|
-
status: fail ? "failed" :
|
|
158
|
-
error: fail ?
|
|
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,
|
|
159
229
|
stepOrder,
|
|
160
230
|
});
|
|
161
|
-
// Let's just not set the id field and have it set only in explicit cases.
|
|
162
|
-
// if (!message.id) {
|
|
163
|
-
// await ctx.db.patch(messageId, {
|
|
164
|
-
// id: messageId,
|
|
165
|
-
// });
|
|
166
|
-
// }
|
|
167
231
|
if (message.fileIds) {
|
|
168
232
|
await changeRefcount(ctx, [], message.fileIds);
|
|
169
233
|
}
|
|
170
234
|
// TODO: delete the associated stream data for the order/stepOrder
|
|
171
235
|
toReturn.push((await ctx.db.get(messageId)));
|
|
172
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
|
+
}
|
|
173
242
|
return { messages: toReturn.map(publicMessage) };
|
|
174
243
|
}
|
|
175
244
|
// exported for tests
|
|
176
245
|
export async function getMaxMessage(ctx, threadId, order) {
|
|
177
|
-
return orderedMessagesStream(ctx,
|
|
246
|
+
return orderedMessagesStream(ctx, {
|
|
247
|
+
threadId,
|
|
248
|
+
sortOrder: "desc",
|
|
249
|
+
startOrder: order,
|
|
250
|
+
startOrderBound: "eq",
|
|
251
|
+
}).first();
|
|
178
252
|
}
|
|
179
|
-
function orderedMessagesStream(ctx,
|
|
253
|
+
function orderedMessagesStream(ctx, args) {
|
|
180
254
|
return mergedStream([true, false].flatMap((tool) => messageStatuses.map((status) => stream(ctx.db, schema)
|
|
181
255
|
.query("messages")
|
|
182
256
|
.withIndex("threadId_status_tool_order_stepOrder", (q) => {
|
|
183
257
|
const qq = q
|
|
184
|
-
.eq("threadId", threadId)
|
|
258
|
+
.eq("threadId", args.threadId)
|
|
185
259
|
.eq("status", status)
|
|
186
260
|
.eq("tool", tool);
|
|
187
|
-
if (
|
|
188
|
-
|
|
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
|
+
}
|
|
189
268
|
}
|
|
190
269
|
return qq;
|
|
191
270
|
})
|
|
192
|
-
.order(sortOrder))), ["order", "stepOrder"]);
|
|
271
|
+
.order(args.sortOrder))), ["order", "stepOrder"]);
|
|
193
272
|
}
|
|
194
|
-
export const
|
|
273
|
+
export const finalizeMessage = mutation({
|
|
195
274
|
args: {
|
|
196
275
|
messageId: v.id("messages"),
|
|
197
|
-
|
|
276
|
+
result: v.union(v.object({ status: v.literal("success") }), v.object({ status: v.literal("failed"), error: v.string() })),
|
|
198
277
|
},
|
|
199
278
|
returns: v.null(),
|
|
200
|
-
handler: async (ctx, { messageId,
|
|
279
|
+
handler: async (ctx, { messageId, result }) => {
|
|
201
280
|
const message = await ctx.db.get(messageId);
|
|
202
281
|
assert(message, `Message ${messageId} not found`);
|
|
203
|
-
|
|
204
|
-
|
|
205
|
-
|
|
206
|
-
|
|
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;
|
|
207
300
|
}
|
|
208
301
|
}
|
|
209
|
-
|
|
210
|
-
|
|
211
|
-
|
|
212
|
-
|
|
213
|
-
|
|
214
|
-
|
|
215
|
-
|
|
216
|
-
|
|
217
|
-
|
|
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
|
+
}
|
|
218
315
|
},
|
|
219
|
-
returns: v.null(),
|
|
220
|
-
handler: commitMessageHandler,
|
|
221
316
|
});
|
|
222
317
|
export const updateMessage = mutation({
|
|
223
318
|
args: {
|
|
224
319
|
messageId: v.id("messages"),
|
|
225
|
-
patch: v.object(
|
|
226
|
-
message
|
|
227
|
-
fileIds
|
|
228
|
-
status
|
|
229
|
-
error
|
|
230
|
-
|
|
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
|
+
]))),
|
|
231
330
|
},
|
|
232
331
|
returns: vMessageDoc,
|
|
233
332
|
handler: async (ctx, args) => {
|
|
@@ -236,83 +335,185 @@ export const updateMessage = mutation({
|
|
|
236
335
|
if (args.patch.fileIds) {
|
|
237
336
|
await changeRefcount(ctx, message.fileIds ?? [], args.patch.fileIds);
|
|
238
337
|
}
|
|
239
|
-
const patch = {
|
|
240
|
-
...args.patch,
|
|
241
|
-
};
|
|
338
|
+
const patch = { ...args.patch };
|
|
242
339
|
if (args.patch.message !== undefined) {
|
|
243
340
|
patch.message = args.patch.message;
|
|
244
341
|
patch.tool = isTool(args.patch.message);
|
|
245
342
|
patch.text = extractText(args.patch.message);
|
|
246
343
|
}
|
|
344
|
+
if (args.patch.status === "failed") {
|
|
345
|
+
if (message.embeddingId) {
|
|
346
|
+
await ctx.db.delete(message.embeddingId);
|
|
347
|
+
}
|
|
348
|
+
patch.embeddingId = undefined;
|
|
349
|
+
}
|
|
247
350
|
await ctx.db.patch(args.messageId, patch);
|
|
248
351
|
return publicMessage((await ctx.db.get(args.messageId)));
|
|
249
352
|
},
|
|
250
353
|
});
|
|
251
|
-
|
|
252
|
-
|
|
253
|
-
|
|
254
|
-
|
|
255
|
-
|
|
256
|
-
|
|
257
|
-
|
|
258
|
-
|
|
259
|
-
|
|
260
|
-
|
|
261
|
-
|
|
262
|
-
|
|
263
|
-
|
|
264
|
-
|
|
265
|
-
}
|
|
266
|
-
export const
|
|
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({
|
|
267
370
|
args: {
|
|
268
|
-
|
|
269
|
-
|
|
270
|
-
/** What order to sort the messages in. To get the latest, use "desc". */
|
|
271
|
-
order: v.union(v.literal("asc"), v.literal("desc")),
|
|
272
|
-
paginationOpts: v.optional(paginationOptsValidator),
|
|
273
|
-
statuses: v.optional(v.array(vMessageStatus)),
|
|
274
|
-
upToAndIncludingMessageId: v.optional(v.id("messages")),
|
|
371
|
+
...cloneMessageArgs,
|
|
372
|
+
paginationOpts: paginationOptsValidator,
|
|
275
373
|
},
|
|
276
374
|
handler: async (ctx, args) => {
|
|
277
|
-
const
|
|
278
|
-
const
|
|
279
|
-
|
|
280
|
-
|
|
281
|
-
|
|
282
|
-
|
|
283
|
-
|
|
284
|
-
.
|
|
285
|
-
|
|
286
|
-
|
|
287
|
-
|
|
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)
|
|
288
390
|
.eq("status", status)
|
|
289
|
-
.eq("tool", tool)
|
|
290
|
-
|
|
291
|
-
|
|
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);
|
|
292
400
|
}
|
|
293
|
-
|
|
294
|
-
|
|
295
|
-
|
|
296
|
-
|
|
297
|
-
|
|
298
|
-
|
|
299
|
-
|
|
300
|
-
|
|
301
|
-
|
|
302
|
-
|
|
303
|
-
|
|
304
|
-
|
|
305
|
-
|
|
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);
|
|
306
477
|
return { ...messages, page: messages.page.map(publicMessage) };
|
|
307
478
|
},
|
|
308
479
|
returns: vPaginationResult(vMessageDoc),
|
|
309
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
|
+
}
|
|
310
513
|
export const getMessagesByIds = query({
|
|
311
|
-
args: {
|
|
312
|
-
messageIds: v.array(v.id("messages")),
|
|
313
|
-
},
|
|
514
|
+
args: { messageIds: v.array(v.id("messages")) },
|
|
314
515
|
handler: async (ctx, args) => {
|
|
315
|
-
return await Promise.all(args.messageIds.map((id) => ctx.db.get(id)));
|
|
516
|
+
return (await Promise.all(args.messageIds.map((id) => ctx.db.get(id)))).map((m) => (m ? publicMessage(m) : null));
|
|
316
517
|
},
|
|
317
518
|
returns: v.array(v.union(v.null(), vMessageDoc)),
|
|
318
519
|
});
|
|
@@ -320,10 +521,12 @@ export const searchMessages = action({
|
|
|
320
521
|
args: {
|
|
321
522
|
threadId: v.optional(v.id("threads")),
|
|
322
523
|
searchAllMessagesForUserId: v.optional(v.string()),
|
|
323
|
-
|
|
524
|
+
targetMessageId: v.optional(v.id("messages")),
|
|
324
525
|
embedding: v.optional(v.array(v.number())),
|
|
325
526
|
embeddingModel: v.optional(v.string()),
|
|
326
527
|
text: v.optional(v.string()),
|
|
528
|
+
textSearch: v.optional(v.boolean()),
|
|
529
|
+
vectorSearch: v.optional(v.boolean()),
|
|
327
530
|
limit: v.number(),
|
|
328
531
|
vectorScoreThreshold: v.optional(v.number()),
|
|
329
532
|
messageRange: v.optional(v.object({ before: v.number(), after: v.number() })),
|
|
@@ -333,23 +536,34 @@ export const searchMessages = action({
|
|
|
333
536
|
assert(args.searchAllMessagesForUserId || args.threadId, "Specify userId or threadId");
|
|
334
537
|
const limit = args.limit;
|
|
335
538
|
let textSearchMessages;
|
|
336
|
-
if (args.
|
|
539
|
+
if (args.textSearch) {
|
|
337
540
|
textSearchMessages = await ctx.runQuery(api.messages.textSearch, {
|
|
338
541
|
searchAllMessagesForUserId: args.searchAllMessagesForUserId,
|
|
339
542
|
threadId: args.threadId,
|
|
543
|
+
targetMessageId: args.targetMessageId,
|
|
340
544
|
text: args.text,
|
|
341
545
|
limit,
|
|
342
|
-
beforeMessageId: args.beforeMessageId,
|
|
343
546
|
});
|
|
344
547
|
}
|
|
345
|
-
if (args.
|
|
346
|
-
|
|
347
|
-
|
|
348
|
-
|
|
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
|
+
}
|
|
349
560
|
}
|
|
350
|
-
|
|
561
|
+
assert(embedding && model, "Embedding missing");
|
|
562
|
+
const dimension = embedding.length;
|
|
563
|
+
validateVectorDimension(dimension);
|
|
564
|
+
const vectors = (await searchVectors(ctx, embedding, {
|
|
351
565
|
dimension,
|
|
352
|
-
model
|
|
566
|
+
model,
|
|
353
567
|
table: "messages",
|
|
354
568
|
searchAllMessagesForUserId: args.searchAllMessagesForUserId,
|
|
355
569
|
threadId: args.threadId,
|
|
@@ -372,7 +586,7 @@ export const searchMessages = action({
|
|
|
372
586
|
embeddingIds,
|
|
373
587
|
textSearchMessages: textSearchMessages?.filter((m) => !embeddingIds.includes(m.embeddingId)),
|
|
374
588
|
messageRange: args.messageRange ?? DEFAULT_MESSAGE_RANGE,
|
|
375
|
-
beforeMessageId: args.
|
|
589
|
+
beforeMessageId: args.targetMessageId,
|
|
376
590
|
limit,
|
|
377
591
|
});
|
|
378
592
|
return messages;
|
|
@@ -414,11 +628,11 @@ export const _fetchSearchMessages = internalQuery({
|
|
|
414
628
|
.map(publicMessage);
|
|
415
629
|
messages.push(...(args.textSearchMessages ?? []));
|
|
416
630
|
// TODO: prioritize more recent messages
|
|
417
|
-
messages
|
|
631
|
+
messages = sorted(messages);
|
|
418
632
|
messages = messages.slice(0, args.limit);
|
|
419
633
|
// Fetch the surrounding messages
|
|
420
634
|
if (!threadId) {
|
|
421
|
-
return messages
|
|
635
|
+
return messages;
|
|
422
636
|
}
|
|
423
637
|
const included = {};
|
|
424
638
|
for (const m of messages) {
|
|
@@ -469,7 +683,7 @@ export const _fetchSearchMessages = internalQuery({
|
|
|
469
683
|
messages.push(publicMessage(r));
|
|
470
684
|
}
|
|
471
685
|
}
|
|
472
|
-
return messages
|
|
686
|
+
return sorted(messages);
|
|
473
687
|
},
|
|
474
688
|
});
|
|
475
689
|
// returns ranges of messages in order of text search relevance,
|
|
@@ -478,21 +692,24 @@ export const textSearch = query({
|
|
|
478
692
|
args: {
|
|
479
693
|
threadId: v.optional(v.id("threads")),
|
|
480
694
|
searchAllMessagesForUserId: v.optional(v.string()),
|
|
481
|
-
text: v.string(),
|
|
695
|
+
text: v.optional(v.string()),
|
|
696
|
+
targetMessageId: v.optional(v.id("messages")),
|
|
482
697
|
limit: v.number(),
|
|
483
|
-
beforeMessageId: v.optional(v.id("messages")),
|
|
484
698
|
},
|
|
485
699
|
handler: async (ctx, args) => {
|
|
486
700
|
assert(args.searchAllMessagesForUserId || args.threadId, "Specify userId or threadId");
|
|
487
|
-
const
|
|
488
|
-
const order =
|
|
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
|
+
}
|
|
489
708
|
const messages = await ctx.db
|
|
490
709
|
.query("messages")
|
|
491
710
|
.withSearchIndex("text_search", (q) => args.searchAllMessagesForUserId
|
|
492
|
-
? q
|
|
493
|
-
|
|
494
|
-
.eq("userId", args.searchAllMessagesForUserId)
|
|
495
|
-
: q.search("text", args.text).eq("threadId", args.threadId))
|
|
711
|
+
? q.search("text", text).eq("userId", args.searchAllMessagesForUserId)
|
|
712
|
+
: q.search("text", text).eq("threadId", args.threadId))
|
|
496
713
|
// Just in case tool messages slip through
|
|
497
714
|
.filter((q) => {
|
|
498
715
|
const qq = q.eq(q.field("tool"), false);
|
|
@@ -503,12 +720,38 @@ export const textSearch = query({
|
|
|
503
720
|
})
|
|
504
721
|
.take(args.limit);
|
|
505
722
|
return messages
|
|
506
|
-
.filter((m) => !
|
|
507
|
-
m.order <
|
|
508
|
-
(m.order ===
|
|
509
|
-
m.stepOrder <
|
|
723
|
+
.filter((m) => !targetMessage ||
|
|
724
|
+
m.order < targetMessage.order ||
|
|
725
|
+
(m.order === targetMessage.order &&
|
|
726
|
+
m.stepOrder < targetMessage.stepOrder))
|
|
510
727
|
.map(publicMessage);
|
|
511
728
|
},
|
|
512
729
|
returns: v.array(vMessageDoc),
|
|
513
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
|
+
});
|
|
514
757
|
//# sourceMappingURL=messages.js.map
|