@convex-dev/agent
Version:
A agent component for Convex.
635 lines (610 loc) • 19.6 kB
text/typescript
import { assert, omit } from "convex-helpers";
import { mergedStream, stream } from "convex-helpers/server/stream";
import { paginationOptsValidator } from "convex/server";
import type { ObjectType } from "convex/values";
import {
DEFAULT_MESSAGE_RANGE,
DEFAULT_RECENT_MESSAGES,
extractText,
isTool,
} from "../shared.js";
import {
vMessageEmbeddings,
vMessageStatus,
vMessageWithMetadataInternal,
vPaginationResult,
vSearchOptions,
} from "../validators.js";
import { api, internal } from "./_generated/api.js";
import type { Doc, Id } from "./_generated/dataModel.js";
import {
action,
internalQuery,
mutation,
type MutationCtx,
query,
type QueryCtx,
} from "./_generated/server.js";
import type { MessageDoc } from "./schema.js";
import { schema, v, vMessageDoc } from "./schema.js";
import {
getThread as _getThread,
listThreadsByUserId as _listThreadsByUserId,
updateThread as _updateThread,
} from "./threads.js";
import { insertVector, searchVectors } from "./vector/index.js";
import {
type VectorDimension,
VectorDimensions,
type VectorTableId,
vVectorId,
} from "./vector/tables.js";
/** @deprecated Use *.threads.listMessagesByThreadId instead. */
export const listThreadsByUserId = _listThreadsByUserId;
/** @deprecated Use *.threads.getThread */
export const getThread = _getThread;
/** @deprecated Use *.threads.updateThread instead */
export const updateThread = _updateThread;
function publicMessage(message: Doc<"messages">): MessageDoc {
return omit(message, ["parentMessageId", "stepId", "files"]);
}
export async function deleteMessage(
ctx: MutationCtx,
messageDoc: Doc<"messages">
) {
await ctx.db.delete(messageDoc._id);
if (messageDoc.embeddingId) {
await ctx.db.delete(messageDoc.embeddingId);
}
for (const fileId of messageDoc.fileIds ?? []) {
if (!fileId) continue;
const file = await ctx.db.get(fileId);
if (file) {
await ctx.db.patch(fileId, { refcount: file.refcount - 1 });
}
}
}
export const messageStatuses = vMessageDoc.fields.status.members.map(
(m) => m.value
);
const addMessagesArgs = {
userId: v.optional(v.string()),
threadId: v.id("threads"),
promptMessageId: v.optional(v.id("messages")),
agentName: v.optional(v.string()),
messages: v.array(vMessageWithMetadataInternal),
embeddings: v.optional(vMessageEmbeddings),
pending: v.optional(v.boolean()),
failPendingSteps: v.optional(v.boolean()),
};
export const addMessages = mutation({
args: addMessagesArgs,
handler: addMessagesHandler,
returns: v.object({
messages: v.array(vMessageDoc),
}),
});
async function addMessagesHandler(
ctx: MutationCtx,
args: ObjectType<typeof addMessagesArgs>
) {
let userId = args.userId;
const threadId = args.threadId;
if (!userId && args.threadId) {
const thread = await ctx.db.get(args.threadId);
assert(thread, `Thread ${args.threadId} not found`);
userId = thread.userId;
}
const {
embeddings,
failPendingSteps,
pending,
messages,
promptMessageId,
...rest
} = args;
const parentMessage = promptMessageId && (await ctx.db.get(promptMessageId));
if (failPendingSteps) {
assert(args.threadId, "threadId is required to fail pending steps");
const pendingMessages = await ctx.db
.query("messages")
.withIndex("threadId_status_tool_order_stepOrder", (q) =>
q.eq("threadId", threadId).eq("status", "pending")
)
.collect();
await Promise.all(
pendingMessages
.filter((m) => !parentMessage || m.order === parentMessage.order)
.map((m) =>
ctx.db.patch(m._id, { status: "failed", error: "Restarting" })
)
);
}
let order, stepOrder;
let fail = false;
if (promptMessageId) {
assert(parentMessage, `Parent message ${promptMessageId} not found`);
if (parentMessage.status === "failed") {
fail = true;
}
order = parentMessage.order;
// Defend against there being existing messages with this parent.
const maxMessage = await getMaxMessage(ctx, threadId, order);
stepOrder = maxMessage?.stepOrder ?? parentMessage.stepOrder;
} else {
const maxMessage = await getMaxMessage(ctx, threadId);
order = maxMessage ? maxMessage.order + 1 : 0;
stepOrder = -1;
}
const toReturn: Doc<"messages">[] = [];
if (embeddings) {
assert(
embeddings.vectors.length === messages.length,
"embeddings.vectors.length must match messages.length"
);
}
for (let i = 0; i < messages.length; i++) {
const message = messages[i];
let embeddingId: VectorTableId | undefined;
if (embeddings && embeddings.vectors[i]) {
embeddingId = await insertVector(ctx, embeddings.dimension, {
vector: embeddings.vectors[i]!,
model: embeddings.model,
table: "messages",
userId,
threadId,
});
}
stepOrder++;
const messageId = await ctx.db.insert("messages", {
...rest,
...message,
embeddingId,
parentMessageId: promptMessageId,
userId,
order,
tool: isTool(message.message),
text: extractText(message.message),
status: fail ? "failed" : pending ? "pending" : "success",
error: fail ? "Parent message failed" : undefined,
stepOrder,
});
// Let's just not set the id field and have it set only in explicit cases.
// if (!message.id) {
// await ctx.db.patch(messageId, {
// id: messageId,
// });
// }
for (const fileId of message.fileIds ?? []) {
if (!fileId) continue;
await ctx.db.patch(fileId, {
refcount: (await ctx.db.get(fileId))!.refcount + 1,
});
}
// TODO: delete the associated stream data for the order/stepOrder
toReturn.push((await ctx.db.get(messageId))!);
}
return { messages: toReturn.map(publicMessage) };
}
// exported for tests
export async function getMaxMessage(
ctx: QueryCtx,
threadId: Id<"threads">,
order?: number
) {
return orderedMessagesStream(ctx, threadId, "desc", order).first();
}
function orderedMessagesStream(
ctx: QueryCtx,
threadId: Id<"threads">,
sortOrder: "asc" | "desc",
order?: number
) {
return mergedStream(
[true, false].flatMap((tool) =>
messageStatuses.map((status) =>
stream(ctx.db, schema)
.query("messages")
.withIndex("threadId_status_tool_order_stepOrder", (q) => {
const qq = q
.eq("threadId", threadId)
.eq("status", status)
.eq("tool", tool);
if (order) {
return qq.eq("order", order);
}
return qq;
})
.order(sortOrder)
)
),
["order", "stepOrder"]
);
}
export const rollbackMessage = mutation({
args: {
messageId: v.id("messages"),
error: v.optional(v.string()),
},
returns: v.null(),
handler: async (ctx, { messageId, error }) => {
const message = await ctx.db.get(messageId);
assert(message, `Message ${messageId} not found`);
const messages = await orderedMessagesStream(
ctx,
message.threadId,
"asc",
message.order
).collect();
for (const m of messages) {
if (m.status === "pending") {
await ctx.db.patch(m._id, { status: "failed", error });
}
}
await ctx.db.patch(messageId, {
status: "failed",
error: error,
});
},
});
export const commitMessage = mutation({
args: {
messageId: v.id("messages"),
},
returns: v.null(),
handler: commitMessageHandler,
});
export const updateMessage = mutation({
args: {
messageId: v.id("messages"),
patch: v.object({
message: v.optional(vMessageDoc.fields.message),
status: v.optional(vMessageStatus),
error: v.optional(v.string()),
}),
},
returns: vMessageDoc,
handler: async (ctx, args) => {
const message = await ctx.db.get(args.messageId);
assert(message, `Message ${args.messageId} not found`);
const patch: Partial<Doc<"messages">> = {
...args.patch,
};
if (args.patch.message !== undefined) {
patch.message = args.patch.message;
patch.tool = isTool(args.patch.message);
patch.text = extractText(args.patch.message);
}
await ctx.db.patch(args.messageId, patch);
return publicMessage((await ctx.db.get(args.messageId))!);
},
});
async function commitMessageHandler(
ctx: MutationCtx,
{ messageId }: { messageId: Id<"messages"> }
) {
const message = await ctx.db.get(messageId);
assert(message, `Message ${messageId} not found`);
const order = message.order!;
const messages = await mergedStream(
[true, false].map((tool) =>
stream(ctx.db, schema)
.query("messages")
.withIndex("threadId_status_tool_order_stepOrder", (q) =>
q
.eq("threadId", message.threadId)
.eq("status", "pending")
.eq("tool", tool)
.eq("order", order)
)
),
["order", "stepOrder"]
).collect();
for (const message of messages) {
await ctx.db.patch(message._id, { status: "success" });
}
}
export const listMessagesByThreadId = query({
args: {
threadId: v.id("threads"),
excludeToolMessages: v.optional(v.boolean()),
/** @deprecated Use excludeToolMessages instead. */
isTool: v.optional(v.literal("use excludeToolMessages instead of this")),
/** What order to sort the messages in. To get the latest, use "desc". */
order: v.union(v.literal("asc"), v.literal("desc")),
paginationOpts: v.optional(paginationOptsValidator),
statuses: v.optional(v.array(vMessageStatus)),
upToAndIncludingMessageId: v.optional(v.id("messages")),
},
handler: async (ctx, args) => {
const statuses =
args.statuses ?? vMessageStatus.members.map((m) => m.value);
const last =
args.upToAndIncludingMessageId &&
(await ctx.db.get(args.upToAndIncludingMessageId));
assert(
!last || last.threadId === args.threadId,
"upToAndIncludingMessageId must be a message in the thread"
);
const toolOptions = args.excludeToolMessages ? [false] : [true, false];
const order = args.order ?? "desc";
const streams = toolOptions.flatMap((tool) =>
statuses.map((status) =>
stream(ctx.db, schema)
.query("messages")
.withIndex("threadId_status_tool_order_stepOrder", (q) => {
const qq = q
.eq("threadId", args.threadId)
.eq("status", status)
.eq("tool", tool);
if (last) {
return qq.lte("order", last.order);
}
return qq;
})
.order(order)
.filterWith(
async (m) =>
!last ||
m.order < last.order ||
(m.order === last.order && m.stepOrder <= last.stepOrder)
)
)
);
const messages = await mergedStream(streams, [
"order",
"stepOrder",
]).paginate(
args.paginationOpts ?? {
numItems: DEFAULT_RECENT_MESSAGES,
cursor: null,
}
);
return { ...messages, page: messages.page.map(publicMessage) };
},
returns: vPaginationResult(vMessageDoc),
});
export const getMessagesByIds = query({
args: {
messageIds: v.array(v.id("messages")),
},
handler: async (ctx, args) => {
return await Promise.all(args.messageIds.map((id) => ctx.db.get(id)));
},
returns: v.array(v.union(v.null(), vMessageDoc)),
});
/** @deprecated Use listMessagesByThreadId instead. */
export const getThreadMessages = query({
args: { deprecated: v.literal("Use listMessagesByThreadId instead") },
handler: async () => {
throw new Error("Use listMessagesByThreadId instead of getThreadMessages");
},
returns: vPaginationResult(vMessageDoc),
});
export const searchMessages = action({
args: {
threadId: v.optional(v.id("threads")),
searchAllMessagesForUserId: v.optional(v.string()),
beforeMessageId: v.optional(v.id("messages")),
...vSearchOptions.fields,
},
returns: v.array(vMessageDoc),
handler: async (ctx, args): Promise<MessageDoc[]> => {
assert(
args.searchAllMessagesForUserId || args.threadId,
"Specify userId or threadId"
);
const limit = args.limit;
let textSearchMessages: MessageDoc[] | undefined;
if (args.text) {
textSearchMessages = await ctx.runQuery(api.messages.textSearch, {
searchAllMessagesForUserId: args.searchAllMessagesForUserId,
threadId: args.threadId,
text: args.text,
limit,
beforeMessageId: args.beforeMessageId,
});
}
if (args.vector) {
const dimension = args.vector.length as VectorDimension;
if (!VectorDimensions.includes(dimension)) {
throw new Error(`Unsupported vector dimension: ${dimension}`);
}
const vectors = (
await searchVectors(ctx, args.vector, {
dimension,
model: args.vectorModel ?? "unknown",
table: "messages",
searchAllMessagesForUserId: args.searchAllMessagesForUserId,
threadId: args.threadId,
limit,
})
).filter((v) => v._score > (args.vectorScoreThreshold ?? 0));
// Reciprocal rank fusion
const k = 10;
const textEmbeddingIds = textSearchMessages?.map((m) => m.embeddingId);
const vectorScores = vectors
.map((v, i) => ({
id: v._id,
score:
1 / (i + k) +
1 / ((textEmbeddingIds?.indexOf(v._id) ?? Infinity) + k),
}))
.sort((a, b) => b.score - a.score);
const vectorIds = vectorScores.slice(0, limit).map((v) => v.id);
const messages: MessageDoc[] = await ctx.runQuery(
internal.messages._fetchSearchMessages,
{
searchAllMessagesForUserId: args.searchAllMessagesForUserId,
threadId: args.threadId,
vectorIds,
textSearchMessages: textSearchMessages?.filter(
(m) => !vectorIds.includes(m.embeddingId! as VectorTableId)
),
messageRange: args.messageRange ?? DEFAULT_MESSAGE_RANGE,
beforeMessageId: args.beforeMessageId,
limit,
}
);
return messages;
}
return textSearchMessages?.flat() ?? [];
},
});
export const _fetchSearchMessages = internalQuery({
args: {
threadId: v.optional(v.id("threads")),
vectorIds: v.array(vVectorId),
searchAllMessagesForUserId: v.optional(v.string()),
textSearchMessages: v.optional(v.array(vMessageDoc)),
messageRange: v.object({ before: v.number(), after: v.number() }),
beforeMessageId: v.optional(v.id("messages")),
limit: v.number(),
},
returns: v.array(vMessageDoc),
handler: async (ctx, args): Promise<MessageDoc[]> => {
const beforeMessage =
args.beforeMessageId && (await ctx.db.get(args.beforeMessageId));
const { searchAllMessagesForUserId, threadId } = args;
assert(
searchAllMessagesForUserId || threadId,
"Specify searchAllMessagesForUserId or threadId to search"
);
let messages: MessageDoc[] = (
await Promise.all(
args.vectorIds.map((embeddingId) =>
ctx.db
.query("messages")
.withIndex("embeddingId", (q) => q.eq("embeddingId", embeddingId))
.filter((q) =>
searchAllMessagesForUserId
? q.eq(q.field("userId"), searchAllMessagesForUserId)
: q.eq(q.field("threadId"), threadId!)
)
// Don't include pending. Failed messages hopefully are deleted but may as well be safe.
.filter((q) => q.eq(q.field("status"), "success"))
.first()
)
)
)
.filter(
(m): m is Doc<"messages"> =>
m !== undefined &&
m !== null &&
!m.tool &&
(!beforeMessage ||
m.order < beforeMessage.order ||
(m.order === beforeMessage.order &&
m.stepOrder < beforeMessage.stepOrder))
)
.map(publicMessage);
messages.push(...(args.textSearchMessages ?? []));
// TODO: prioritize more recent messages
messages.sort((a, b) => a.order! - b.order!);
messages = messages.slice(0, args.limit);
// Fetch the surrounding messages
if (!threadId) {
return messages.sort((a, b) => a.order - b.order);
}
const included: Record<string, Set<number>> = {};
for (const m of messages) {
const searchId = m.threadId ?? m.userId!;
if (!included[searchId]) {
included[searchId] = new Set();
}
included[searchId].add(m.order!);
}
const ranges: Record<string, Doc<"messages">[]> = {};
const { before, after } = args.messageRange;
for (const m of messages) {
const searchId = m.threadId ?? m.userId!;
const order = m.order!;
let earliest = order - before;
let latest = order + after;
for (; earliest <= latest; earliest++) {
if (!included[searchId].has(earliest)) {
break;
}
}
for (; latest >= earliest; latest--) {
if (!included[searchId].has(latest)) {
break;
}
}
for (let i = earliest; i <= latest; i++) {
included[searchId].add(i);
}
if (earliest !== latest) {
const surrounding = await ctx.db
.query("messages")
.withIndex("threadId_status_tool_order_stepOrder", (q) =>
q
.eq("threadId", m.threadId as Id<"threads">)
.eq("status", "success")
.eq("tool", false)
.gte("order", earliest)
.lte("order", latest)
)
.collect();
if (!ranges[searchId]) {
ranges[searchId] = [];
}
ranges[searchId].push(...surrounding);
}
}
for (const r of Object.values(ranges).flat()) {
if (!messages.some((m) => m._id === r._id)) {
messages.push(publicMessage(r));
}
}
return messages.sort((a, b) => a.order - b.order);
},
});
// returns ranges of messages in order of text search relevance,
// excluding duplicates in later ranges.
export const textSearch = query({
args: {
threadId: v.optional(v.id("threads")),
searchAllMessagesForUserId: v.optional(v.string()),
text: v.string(),
limit: v.number(),
beforeMessageId: v.optional(v.id("messages")),
},
handler: async (ctx, args) => {
assert(
args.searchAllMessagesForUserId || args.threadId,
"Specify userId or threadId"
);
const beforeMessage =
args.beforeMessageId && (await ctx.db.get(args.beforeMessageId));
const order = beforeMessage?.order;
const messages = await ctx.db
.query("messages")
.withSearchIndex("text_search", (q) =>
args.searchAllMessagesForUserId
? q
.search("text", args.text)
.eq("userId", args.searchAllMessagesForUserId)
: q.search("text", args.text).eq("threadId", args.threadId!)
)
// Just in case tool messages slip through
.filter((q) => {
const qq = q.eq(q.field("tool"), false);
if (order) {
return q.and(qq, q.lte(q.field("order"), order));
}
return qq;
})
.take(args.limit);
return messages
.filter(
(m) =>
!beforeMessage ||
m.order < beforeMessage.order ||
(m.order === beforeMessage.order &&
m.stepOrder < beforeMessage.stepOrder)
)
.map(publicMessage);
},
returns: v.array(vMessageDoc),
});