@convex-dev/agent 0.0.17-alpha.2 → 0.1.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.
Files changed (49) hide show
  1. package/README.md +25 -7
  2. package/dist/commonjs/client/index.d.ts +229 -756
  3. package/dist/commonjs/client/index.d.ts.map +1 -1
  4. package/dist/commonjs/client/index.js +135 -140
  5. package/dist/commonjs/client/index.js.map +1 -1
  6. package/dist/commonjs/component/messages.d.ts +208 -896
  7. package/dist/commonjs/component/messages.d.ts.map +1 -1
  8. package/dist/commonjs/component/messages.js +117 -109
  9. package/dist/commonjs/component/messages.js.map +1 -1
  10. package/dist/commonjs/component/schema.d.ts +1278 -580
  11. package/dist/commonjs/component/schema.d.ts.map +1 -1
  12. package/dist/commonjs/component/schema.js +28 -21
  13. package/dist/commonjs/component/schema.js.map +1 -1
  14. package/dist/commonjs/component/users.d.ts +5 -2
  15. package/dist/commonjs/component/users.d.ts.map +1 -1
  16. package/dist/commonjs/component/users.js +57 -27
  17. package/dist/commonjs/component/users.js.map +1 -1
  18. package/dist/commonjs/validators.d.ts +44 -59
  19. package/dist/commonjs/validators.d.ts.map +1 -1
  20. package/dist/commonjs/validators.js +10 -12
  21. package/dist/commonjs/validators.js.map +1 -1
  22. package/dist/esm/client/index.d.ts +229 -756
  23. package/dist/esm/client/index.d.ts.map +1 -1
  24. package/dist/esm/client/index.js +135 -140
  25. package/dist/esm/client/index.js.map +1 -1
  26. package/dist/esm/component/messages.d.ts +208 -896
  27. package/dist/esm/component/messages.d.ts.map +1 -1
  28. package/dist/esm/component/messages.js +117 -109
  29. package/dist/esm/component/messages.js.map +1 -1
  30. package/dist/esm/component/schema.d.ts +1278 -580
  31. package/dist/esm/component/schema.d.ts.map +1 -1
  32. package/dist/esm/component/schema.js +28 -21
  33. package/dist/esm/component/schema.js.map +1 -1
  34. package/dist/esm/component/users.d.ts +5 -2
  35. package/dist/esm/component/users.d.ts.map +1 -1
  36. package/dist/esm/component/users.js +57 -27
  37. package/dist/esm/component/users.js.map +1 -1
  38. package/dist/esm/validators.d.ts +44 -59
  39. package/dist/esm/validators.d.ts.map +1 -1
  40. package/dist/esm/validators.js +10 -12
  41. package/dist/esm/validators.js.map +1 -1
  42. package/package.json +2 -2
  43. package/src/client/index.ts +215 -199
  44. package/src/component/_generated/api.d.ts +19 -94
  45. package/src/component/messages.test.ts +110 -3
  46. package/src/component/messages.ts +176 -154
  47. package/src/component/schema.ts +34 -21
  48. package/src/component/users.ts +57 -32
  49. package/src/validators.ts +12 -15
@@ -1,4 +1,4 @@
1
- import { assert } from "convex-helpers";
1
+ import { assert, omit } from "convex-helpers";
2
2
  import { mergedStream, stream } from "convex-helpers/server/stream";
3
3
  import { ObjectType } from "convex/values";
4
4
  import {
@@ -9,6 +9,7 @@ import {
9
9
  } from "../shared.js";
10
10
  import {
11
11
  paginationResultValidator,
12
+ vMessageEmbeddings,
12
13
  vMessageStatus,
13
14
  vMessageWithMetadata,
14
15
  vSearchOptions,
@@ -38,16 +39,20 @@ import {
38
39
  updateThread as _updateThread,
39
40
  } from "./threads.js";
40
41
  import { paginationOptsValidator } from "convex/server";
41
-
42
+ import { MessageDoc, vMessageDoc } from "./schema.js";
42
43
 
43
44
  /** @deprecated Use *.threads.listMessagesByThreadId instead. */
44
- export const listThreadsByUserId= _listThreadsByUserId
45
+ export const listThreadsByUserId = _listThreadsByUserId;
45
46
 
46
47
  /** @deprecated Use *.threads.getThread */
47
48
  export const getThread = _getThread;
48
49
 
49
50
  /** @deprecated Use *.threads.updateThread instead */
50
- export const updateThread= _updateThread;
51
+ export const updateThread = _updateThread;
52
+
53
+ function publicMessage(message: Doc<"messages">): MessageDoc {
54
+ return omit(message, ["parentMessageId", "stepId"]);
55
+ }
51
56
 
52
57
  export async function deleteMessage(
53
58
  ctx: MutationCtx,
@@ -66,7 +71,6 @@ export async function deleteMessage(
66
71
  }
67
72
  }
68
73
 
69
- export const vMessageDoc = schema.tables.messages.validator;
70
74
  export const messageStatuses = vMessageDoc.fields.status.members.map(
71
75
  (m) => m.value
72
76
  );
@@ -75,9 +79,10 @@ const addMessagesArgs = {
75
79
  userId: v.optional(v.string()),
76
80
  threadId: v.id("threads"),
77
81
  stepId: v.optional(v.id("steps")),
78
- parentMessageId: v.optional(v.id("messages")),
82
+ promptMessageId: v.optional(v.id("messages")),
79
83
  agentName: v.optional(v.string()),
80
84
  messages: v.array(vMessageWithMetadata),
85
+ embeddings: v.optional(vMessageEmbeddings),
81
86
  pending: v.optional(v.boolean()),
82
87
  failPendingSteps: v.optional(v.boolean()),
83
88
  };
@@ -85,8 +90,8 @@ export const addMessages = mutation({
85
90
  args: addMessagesArgs,
86
91
  handler: addMessagesHandler,
87
92
  returns: v.object({
88
- messages: v.array(v.doc("messages")),
89
- pending: v.optional(v.doc("messages")),
93
+ messages: v.array(vMessageDoc),
94
+ pending: v.optional(vMessageDoc),
90
95
  }),
91
96
  });
92
97
  async function addMessagesHandler(
@@ -100,8 +105,15 @@ async function addMessagesHandler(
100
105
  assert(thread, `Thread ${args.threadId} not found`);
101
106
  userId = thread.userId;
102
107
  }
103
- const { failPendingSteps, pending, messages, parentMessageId, ...rest } =
104
- args;
108
+ const {
109
+ embeddings,
110
+ failPendingSteps,
111
+ pending,
112
+ messages,
113
+ promptMessageId,
114
+ ...rest
115
+ } = args;
116
+ const parentMessage = promptMessageId && (await ctx.db.get(promptMessageId));
105
117
  if (failPendingSteps) {
106
118
  assert(args.threadId, "threadId is required to fail pending steps");
107
119
  const pendingMessages = await ctx.db
@@ -111,57 +123,69 @@ async function addMessagesHandler(
111
123
  )
112
124
  .collect();
113
125
  await Promise.all(
114
- pendingMessages.map((m) =>
115
- ctx.db.patch(m._id, { status: "failed", error: "Restarting" })
116
- )
126
+ pendingMessages
127
+ .filter((m) => !parentMessage || m.order === parentMessage.order)
128
+ .map((m) =>
129
+ ctx.db.patch(m._id, { status: "failed", error: "Restarting" })
130
+ )
117
131
  );
118
132
  }
119
- const maxMessage = await getMaxMessage(ctx, threadId, userId);
120
- let order = maxMessage?.order ?? -1;
121
- let stepOrder = maxMessage?.stepOrder ?? 0;
122
- let lastMessageIsTool = maxMessage?.tool ?? false;
133
+ let order, stepOrder;
134
+ let fail = false;
135
+ if (promptMessageId) {
136
+ assert(parentMessage, `Parent message ${promptMessageId} not found`);
137
+ if (parentMessage.status === "failed") {
138
+ fail = true;
139
+ }
140
+ order = parentMessage.order;
141
+ // Defend against there being existing messages with this parent.
142
+ const maxMessage = await getMaxMessage(ctx, threadId, order);
143
+ stepOrder = maxMessage?.stepOrder ?? parentMessage.stepOrder;
144
+ } else {
145
+ const maxMessage = await getMaxMessage(ctx, threadId);
146
+ order = maxMessage ? maxMessage.order + 1 : 0;
147
+ stepOrder = -1;
148
+ }
123
149
  const toReturn: Doc<"messages">[] = [];
124
150
  if (messages.length > 0) {
125
- for (const { message, files, embedding, ...fields } of messages) {
151
+ if (embeddings) {
152
+ assert(
153
+ embeddings.vectors.length === messages.length,
154
+ "embeddings.vectors.length must match messages.length"
155
+ );
156
+ }
157
+ for (let i = 0; i < messages.length; i++) {
158
+ const message = messages[i];
126
159
  let embeddingId: VectorTableId | undefined;
127
- if (embedding) {
128
- embeddingId = await insertVector(ctx, embedding.dimension, {
129
- vector: embedding.vector,
130
- model: embedding.model,
160
+ if (embeddings && embeddings.vectors[i]) {
161
+ embeddingId = await insertVector(ctx, embeddings.dimension, {
162
+ vector: embeddings.vectors[i]!,
163
+ model: embeddings.model,
131
164
  table: "messages",
132
165
  userId,
133
166
  threadId,
134
167
  });
135
168
  }
136
- const tool = isTool(message);
137
- if (lastMessageIsTool) {
138
- stepOrder++;
139
- } else {
140
- order++;
141
- stepOrder = 0;
142
- }
143
- lastMessageIsTool = tool;
144
- const text = extractText(message);
169
+ stepOrder++;
145
170
  const messageId = await ctx.db.insert("messages", {
146
171
  ...rest,
147
- ...fields,
172
+ ...message,
148
173
  embeddingId,
149
- parentMessageId,
174
+ parentMessageId: promptMessageId,
150
175
  userId,
151
- message,
152
176
  order,
153
- tool,
154
- text,
155
- files,
156
- status: pending ? "pending" : "success",
177
+ tool: isTool(message.message),
178
+ text: extractText(message.message),
179
+ status: fail ? "failed" : pending ? "pending" : "success",
180
+ error: fail ? "Parent message failed" : undefined,
157
181
  stepOrder,
158
182
  });
159
- if (!fields.id) {
183
+ if (!message.id) {
160
184
  await ctx.db.patch(messageId, {
161
185
  id: messageId,
162
186
  });
163
187
  }
164
- for (const { fileId } of files ?? []) {
188
+ for (const { fileId } of message.files ?? []) {
165
189
  if (!fileId) continue;
166
190
  await ctx.db.patch(fileId, {
167
191
  refcount: (await ctx.db.get(fileId))!.refcount + 1,
@@ -176,45 +200,44 @@ async function addMessagesHandler(
176
200
  // exported for tests
177
201
  export async function getMaxMessage(
178
202
  ctx: QueryCtx,
179
- threadId: Id<"threads"> | undefined,
180
- userId: string | undefined
203
+ threadId: Id<"threads">,
204
+ order?: number
181
205
  ) {
182
- assert(threadId || userId, "One of threadId or userId is required");
183
- if (threadId) {
184
- return mergedStream(
185
- [true, false].flatMap((tool) =>
186
- ["success" as const, "pending" as const].map((status) =>
187
- stream(ctx.db, schema)
188
- .query("messages")
189
- .withIndex("threadId_status_tool_order_stepOrder", (q) =>
190
- q.eq("threadId", threadId).eq("status", status).eq("tool", tool)
191
- )
192
- .order("desc")
193
- )
194
- ),
195
- ["order", "stepOrder"]
196
- ).first();
197
- } else {
198
- return mergedStream(
199
- [true, false].flatMap((tool) =>
200
- ["success" as const, "pending" as const].map((status) =>
201
- stream(ctx.db, schema)
202
- .query("messages")
203
- .withIndex("userId_status_tool_order_stepOrder", (q) =>
204
- q.eq("userId", userId).eq("status", status).eq("tool", tool)
205
- )
206
- .order("desc")
207
- )
208
- ),
209
- ["order", "stepOrder"]
210
- ).first();
211
- }
206
+ return orderedMessagesStream(ctx, threadId, "desc", order).first();
207
+ }
208
+
209
+ function orderedMessagesStream(
210
+ ctx: QueryCtx,
211
+ threadId: Id<"threads">,
212
+ sortOrder: "asc" | "desc",
213
+ order?: number
214
+ ) {
215
+ return mergedStream(
216
+ [true, false].flatMap((tool) =>
217
+ ["success" as const, "pending" as const].map((status) =>
218
+ stream(ctx.db, schema)
219
+ .query("messages")
220
+ .withIndex("threadId_status_tool_order_stepOrder", (q) => {
221
+ const qq = q
222
+ .eq("threadId", threadId)
223
+ .eq("status", status)
224
+ .eq("tool", tool);
225
+ if (order) {
226
+ return qq.eq("order", order);
227
+ }
228
+ return qq;
229
+ })
230
+ .order(sortOrder)
231
+ )
232
+ ),
233
+ ["order", "stepOrder"]
234
+ );
212
235
  }
213
236
 
214
237
  const addStepArgs = {
215
238
  userId: v.optional(v.string()),
216
239
  threadId: v.id("threads"),
217
- parentMessageId: v.id("messages"),
240
+ promptMessageId: v.id("messages"),
218
241
  step: vStepWithMessages,
219
242
  failPendingSteps: v.optional(v.boolean()),
220
243
  };
@@ -228,16 +251,15 @@ async function addStepHandler(
228
251
  ctx: MutationCtx,
229
252
  args: ObjectType<typeof addStepArgs>
230
253
  ) {
231
- const parentMessage = await ctx.db.get(args.parentMessageId);
232
- assert(parentMessage, `Message ${args.parentMessageId} not found`);
254
+ const parentMessage = await ctx.db.get(args.promptMessageId);
255
+ assert(parentMessage, `Message ${args.promptMessageId} not found`);
233
256
  const order = parentMessage.order;
234
- assert(order !== undefined, `${args.parentMessageId} has no order`);
235
257
  // TODO: only fetch the last one if we aren't failing pending steps
236
258
  let steps = await ctx.db
237
259
  .query("steps")
238
260
  .withIndex("parentMessageId_order_stepOrder", (q) =>
239
261
  // TODO: fetch pending, and commit later
240
- q.eq("parentMessageId", args.parentMessageId)
262
+ q.eq("parentMessageId", args.promptMessageId)
241
263
  )
242
264
  .collect();
243
265
  if (args.failPendingSteps) {
@@ -251,7 +273,7 @@ async function addStepHandler(
251
273
  const { step, messages } = args.step;
252
274
  const stepId = await ctx.db.insert("steps", {
253
275
  threadId: args.threadId,
254
- parentMessageId: args.parentMessageId,
276
+ parentMessageId: args.promptMessageId,
255
277
  order,
256
278
  stepOrder: (steps.at(-1)?.stepOrder ?? -1) + 1,
257
279
  status: step.finishReason === "stop" ? "success" : "pending",
@@ -261,7 +283,7 @@ async function addStepHandler(
261
283
  userId: args.userId,
262
284
  threadId: args.threadId,
263
285
  stepId,
264
- parentMessageId: args.parentMessageId,
286
+ promptMessageId: args.promptMessageId,
265
287
  agentName: parentMessage.agentName,
266
288
  messages,
267
289
  pending: step.finishReason === "stop" ? false : true,
@@ -269,7 +291,7 @@ async function addStepHandler(
269
291
  });
270
292
  // We don't commit if the parent is still pending.
271
293
  if (step.finishReason === "stop") {
272
- await commitMessageHandler(ctx, { messageId: args.parentMessageId });
294
+ await commitMessageHandler(ctx, { messageId: args.promptMessageId });
273
295
  }
274
296
  steps.push((await ctx.db.get(stepId))!);
275
297
  return steps;
@@ -284,8 +306,18 @@ export const rollbackMessage = mutation({
284
306
  handler: async (ctx, { messageId, error }) => {
285
307
  const message = await ctx.db.get(messageId);
286
308
  assert(message, `Message ${messageId} not found`);
287
- // TODO: do BFS to fail all associated messages, then steps
288
- // with parentMessageId of those messages, etc.
309
+ const messages = await orderedMessagesStream(
310
+ ctx,
311
+ message.threadId,
312
+ "asc",
313
+ message.order
314
+ ).collect();
315
+ for (const m of messages) {
316
+ if (m.status === "pending") {
317
+ await ctx.db.patch(m._id, { status: "failed", error });
318
+ }
319
+ }
320
+
289
321
  const steps = await ctx.db
290
322
  .query("steps")
291
323
  .withIndex("parentMessageId_order_stepOrder", (q) =>
@@ -347,7 +379,6 @@ async function commitMessageHandler(
347
379
  ).collect();
348
380
  for (const message of messages) {
349
381
  await ctx.db.patch(message._id, { status: "success" });
350
- // TODO: recursively commit steps & messages that might depend on this one.
351
382
  }
352
383
  }
353
384
 
@@ -360,13 +391,18 @@ export const listMessagesByThreadId = query({
360
391
  order: v.optional(v.union(v.literal("asc"), v.literal("desc"))),
361
392
  paginationOpts: v.optional(paginationOptsValidator),
362
393
  statuses: v.optional(v.array(vMessageStatus)),
363
- beforeMessageId: v.optional(v.id("messages")),
394
+ upToAndIncludingMessageId: v.optional(v.id("messages")),
364
395
  },
365
396
  handler: async (ctx, args) => {
366
397
  const statuses =
367
398
  args.statuses ?? vMessageStatus.members.map((m) => m.value);
368
- const before =
369
- args.beforeMessageId && (await ctx.db.get(args.beforeMessageId));
399
+ const last =
400
+ args.upToAndIncludingMessageId &&
401
+ (await ctx.db.get(args.upToAndIncludingMessageId));
402
+ assert(
403
+ !last || last.threadId === args.threadId,
404
+ "upToAndIncludingMessageId must be a message in the thread"
405
+ );
370
406
  const toolOptions = args.excludeToolMessages ? [false] : [true, false];
371
407
  const order = args.order ?? "desc";
372
408
  const streams = toolOptions.flatMap((tool) =>
@@ -378,17 +414,17 @@ export const listMessagesByThreadId = query({
378
414
  .eq("threadId", args.threadId)
379
415
  .eq("status", status)
380
416
  .eq("tool", tool);
381
- if (before) {
382
- return qq.lte("order", before.order);
417
+ if (last) {
418
+ return qq.lte("order", last.order);
383
419
  }
384
420
  return qq;
385
421
  })
386
422
  .order(order)
387
423
  .filterWith(
388
424
  async (m) =>
389
- !before ||
390
- m.order < before.order ||
391
- (m.order === before.order && m.stepOrder < before.stepOrder)
425
+ !last ||
426
+ m.order < last.order ||
427
+ (m.order === last.order && m.stepOrder <= last.stepOrder)
392
428
  )
393
429
  )
394
430
  );
@@ -401,9 +437,9 @@ export const listMessagesByThreadId = query({
401
437
  cursor: null,
402
438
  }
403
439
  );
404
- return messages;
440
+ return { ...messages, page: messages.page.map(publicMessage) };
405
441
  },
406
- returns: paginationResultValidator(v.doc("messages")),
442
+ returns: paginationResultValidator(vMessageDoc),
407
443
  });
408
444
 
409
445
  /** @deprecated Use listMessagesByThreadId instead. */
@@ -412,7 +448,7 @@ export const getThreadMessages = query({
412
448
  handler: async () => {
413
449
  throw new Error("Use listMessagesByThreadId instead of getThreadMessages");
414
450
  },
415
- returns: paginationResultValidator(v.doc("messages")),
451
+ returns: paginationResultValidator(vMessageDoc),
416
452
  });
417
453
 
418
454
  export const searchMessages = action({
@@ -422,11 +458,11 @@ export const searchMessages = action({
422
458
  beforeMessageId: v.optional(v.id("messages")),
423
459
  ...vSearchOptions.fields,
424
460
  },
425
- returns: v.array(v.doc("messages")),
426
- handler: async (ctx, args): Promise<Doc<"messages">[]> => {
461
+ returns: v.array(vMessageDoc),
462
+ handler: async (ctx, args): Promise<MessageDoc[]> => {
427
463
  assert(args.userId || args.threadId, "Specify userId or threadId");
428
464
  const limit = args.limit;
429
- let textSearchMessages: Doc<"messages">[] | undefined;
465
+ let textSearchMessages: MessageDoc[] | undefined;
430
466
  if (args.text) {
431
467
  textSearchMessages = await ctx.runQuery(api.messages.textSearch, {
432
468
  userId: args.userId,
@@ -463,14 +499,14 @@ export const searchMessages = action({
463
499
  }))
464
500
  .sort((a, b) => b.score - a.score);
465
501
  const vectorIds = vectorScores.slice(0, limit).map((v) => v.id);
466
- const messages: Doc<"messages">[] = await ctx.runQuery(
502
+ const messages: MessageDoc[] = await ctx.runQuery(
467
503
  internal.messages._fetchSearchMessages,
468
504
  {
469
505
  userId: args.userId,
470
506
  threadId: args.threadId,
471
507
  vectorIds,
472
508
  textSearchMessages: textSearchMessages?.filter(
473
- (m) => !vectorIds.includes(m.embeddingId!)
509
+ (m) => !vectorIds.includes(m.embeddingId! as VectorTableId)
474
510
  ),
475
511
  messageRange: args.messageRange ?? DEFAULT_MESSAGE_RANGE,
476
512
  beforeMessageId: args.beforeMessageId,
@@ -488,18 +524,18 @@ export const _fetchSearchMessages = internalQuery({
488
524
  userId: v.optional(v.string()),
489
525
  threadId: v.optional(v.id("threads")),
490
526
  vectorIds: v.array(vVectorId),
491
- textSearchMessages: v.optional(v.array(v.doc("messages"))),
527
+ textSearchMessages: v.optional(v.array(vMessageDoc)),
492
528
  messageRange: v.object({ before: v.number(), after: v.number() }),
493
529
  beforeMessageId: v.optional(v.id("messages")),
494
530
  limit: v.number(),
495
531
  },
496
- returns: v.array(v.doc("messages")),
497
- handler: async (ctx, args): Promise<Doc<"messages">[]> => {
532
+ returns: v.array(vMessageDoc),
533
+ handler: async (ctx, args): Promise<MessageDoc[]> => {
498
534
  const beforeMessage =
499
535
  args.beforeMessageId && (await ctx.db.get(args.beforeMessageId));
500
536
  const { userId, threadId } = args;
501
537
  assert(userId || threadId, "Specify userId or threadId to search");
502
- let messages = (
538
+ let messages: MessageDoc[] = (
503
539
  await Promise.all(
504
540
  args.vectorIds.map((embeddingId) =>
505
541
  ctx.db
@@ -515,16 +551,18 @@ export const _fetchSearchMessages = internalQuery({
515
551
  .first()
516
552
  )
517
553
  )
518
- ).filter(
519
- (m): m is Doc<"messages"> =>
520
- m !== undefined &&
521
- m !== null &&
522
- !m.tool &&
523
- (!beforeMessage ||
524
- m.order < beforeMessage.order ||
525
- (m.order === beforeMessage.order &&
526
- m.stepOrder < beforeMessage.stepOrder))
527
- );
554
+ )
555
+ .filter(
556
+ (m): m is Doc<"messages"> =>
557
+ m !== undefined &&
558
+ m !== null &&
559
+ !m.tool &&
560
+ (!beforeMessage ||
561
+ m.order < beforeMessage.order ||
562
+ (m.order === beforeMessage.order &&
563
+ m.stepOrder < beforeMessage.stepOrder))
564
+ )
565
+ .map(publicMessage);
528
566
  messages.push(...(args.textSearchMessages ?? []));
529
567
  // TODO: prioritize more recent messages
530
568
  messages.sort((a, b) => a.order! - b.order!);
@@ -562,44 +600,26 @@ export const _fetchSearchMessages = internalQuery({
562
600
  included[searchId].add(i);
563
601
  }
564
602
  if (earliest !== latest) {
565
- if (m.threadId) {
566
- const surrounding = await ctx.db
567
- .query("messages")
568
- .withIndex("threadId_status_tool_order_stepOrder", (q) =>
569
- q
570
- .eq("threadId", m.threadId)
571
- .eq("status", "success")
572
- .eq("tool", false)
573
- .gte("order", earliest)
574
- .lte("order", latest)
575
- )
576
- .collect();
577
- if (!ranges[searchId]) {
578
- ranges[searchId] = [];
579
- }
580
- ranges[searchId].push(...surrounding);
581
- } else {
582
- const surrounding = await ctx.db
583
- .query("messages")
584
- .withIndex("userId_status_tool_order_stepOrder", (q) =>
585
- q
586
- .eq("userId", m.userId!)
587
- .eq("status", "success")
588
- .eq("tool", false)
589
- .gte("order", earliest)
590
- .lte("order", latest)
591
- )
592
- .collect();
593
- if (!ranges[searchId]) {
594
- ranges[searchId] = [];
595
- }
596
- ranges[searchId].push(...surrounding);
603
+ const surrounding = await ctx.db
604
+ .query("messages")
605
+ .withIndex("threadId_status_tool_order_stepOrder", (q) =>
606
+ q
607
+ .eq("threadId", m.threadId as Id<"threads">)
608
+ .eq("status", "success")
609
+ .eq("tool", false)
610
+ .gte("order", earliest)
611
+ .lte("order", latest)
612
+ )
613
+ .collect();
614
+ if (!ranges[searchId]) {
615
+ ranges[searchId] = [];
597
616
  }
617
+ ranges[searchId].push(...surrounding);
598
618
  }
599
619
  }
600
620
  for (const r of Object.values(ranges).flat()) {
601
621
  if (!messages.some((m) => m._id === r._id)) {
602
- messages.push(r);
622
+ messages.push(publicMessage(r));
603
623
  }
604
624
  }
605
625
  return messages.sort((a, b) => a.order - b.order);
@@ -637,13 +657,15 @@ export const textSearch = query({
637
657
  return qq;
638
658
  })
639
659
  .take(args.limit);
640
- return messages.filter(
641
- (m) =>
642
- !beforeMessage ||
643
- m.order < beforeMessage.order ||
644
- (m.order === beforeMessage.order &&
645
- m.stepOrder < beforeMessage.stepOrder)
646
- );
660
+ return messages
661
+ .filter(
662
+ (m) =>
663
+ !beforeMessage ||
664
+ m.order < beforeMessage.order ||
665
+ (m.order === beforeMessage.order &&
666
+ m.stepOrder < beforeMessage.stepOrder)
667
+ )
668
+ .map(publicMessage);
647
669
  },
648
- returns: v.array(v.doc("messages")),
670
+ returns: v.array(vMessageDoc),
649
671
  });
@@ -1,5 +1,5 @@
1
1
  import { defineSchema, defineTable } from "convex/server";
2
- import { v } from "convex/values";
2
+ import { Infer, v } from "convex/values";
3
3
  import {
4
4
  vThreadStatus,
5
5
  vMessage,
@@ -13,9 +13,11 @@ import {
13
13
  vProviderMetadata,
14
14
  vReasoningDetails,
15
15
  vFile,
16
+ vFileWithStringId,
16
17
  } from "../validators.js";
17
18
  import { typedV } from "convex-helpers/validators";
18
19
  import vectorTables, { vVectorId } from "./vector/tables.js";
20
+ import { omit } from "convex-helpers";
19
21
 
20
22
  export const schema = defineSchema({
21
23
  threads: defineTable({
@@ -28,14 +30,10 @@ export const schema = defineSchema({
28
30
  parentThreadIds: v.optional(v.array(v.id("threads"))),
29
31
  order: /*DEPRECATED*/ v.optional(v.number()),
30
32
  }).index("userId", ["userId"]),
31
- // TODO: text search on title/ summary
32
33
  messages: defineTable({
33
34
  id: v.optional(v.string()), // external id, e.g. from Vercel AI SDK
34
- userId: v.optional(v.string()), // useful for future indexes (text search)
35
+ userId: v.optional(v.string()), // useful for searching across threads
35
36
  threadId: v.id("threads"),
36
- // TODO: is this redunant with message at last step @ order - 1?
37
- parentMessageId: v.optional(v.id("messages")),
38
- stepId: v.optional(v.id("steps")),
39
37
  // Repeats until a non-tool message.
40
38
  order: v.number(),
41
39
  stepOrder: v.number(),
@@ -64,6 +62,9 @@ export const schema = defineSchema({
64
62
  reasoningDetails: v.optional(vReasoningDetails),
65
63
  warnings: v.optional(v.array(vLanguageModelV1CallWarning)),
66
64
  finishReason: v.optional(vFinishReason),
65
+ // DEPRECATED
66
+ parentMessageId: v.optional(v.id("messages")),
67
+ stepId: v.optional(v.id("steps")),
67
68
  })
68
69
  // Allows finding successful visible messages in order
69
70
  // Also surface pending messages separately to e.g. stream
@@ -75,21 +76,6 @@ export const schema = defineSchema({
75
76
  "order",
76
77
  "stepOrder",
77
78
  ])
78
- .index("userId_status_tool_order_stepOrder", [
79
- "userId",
80
- "status",
81
- "tool",
82
- "order",
83
- "stepOrder",
84
- ])
85
- // Allows finding all threaded messages in order
86
- // Allows finding all failed messages to evaluate
87
- // .index("status_parentMessageId_order_stepOrder", [
88
- // "status",
89
- // "parentMessageId",
90
- // "order",
91
- // "stepOrder",
92
- // ])
93
79
  // Allows text search on message content
94
80
  .searchIndex("text_search", {
95
81
  searchField: "text",
@@ -148,4 +134,31 @@ export const schema = defineSchema({
148
134
  export const vv = typedV(schema);
149
135
  export { vv as v };
150
136
 
137
+ // Public
138
+ export const vThreadDoc = v.object({
139
+ _id: v.string(),
140
+ _creationTime: v.number(),
141
+ userId: v.optional(v.string()), // Unset for anonymous
142
+ title: v.optional(v.string()),
143
+ summary: v.optional(v.string()),
144
+ status: vThreadStatus,
145
+ });
146
+ export type ThreadDoc = Infer<typeof vThreadDoc>;
147
+
148
+ export const vMessageDoc = v.object({
149
+ _id: v.string(),
150
+ _creationTime: v.number(),
151
+ ...omit(schema.tables.messages.validator.fields, [
152
+ "parentMessageId",
153
+ "stepId",
154
+ ]),
155
+ // Overwrite all the types that have a v.id validator
156
+ // Outside of the component, they are strings
157
+ threadId: v.string(),
158
+ embeddingId: v.optional(v.string()),
159
+ files: v.optional(v.array(vFileWithStringId)),
160
+ });
161
+ export type MessageDoc = Infer<typeof vMessageDoc>;
162
+
163
+
151
164
  export default schema;