UNPKG

@convex-dev/agent

Version:

A agent component for Convex.

478 lines 19.8 kB
import { assert, omit } from "convex-helpers"; import { mergedStream, stream } from "convex-helpers/server/stream"; import { paginationOptsValidator } from "convex/server"; 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 { action, internalQuery, mutation, query, } from "./_generated/server.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 { VectorDimensions, 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) { return omit(message, ["parentMessageId", "stepId", "files"]); } export async function deleteMessage(ctx, messageDoc) { 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, args) { 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 = []; 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; 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, threadId, order) { return orderedMessagesStream(ctx, threadId, "desc", order).first(); } function orderedMessagesStream(ctx, threadId, sortOrder, order) { 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 = { ...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, { messageId }) { 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) => { assert(args.searchAllMessagesForUserId || args.threadId, "Specify userId or threadId"); const limit = args.limit; let textSearchMessages; 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; 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 = await ctx.runQuery(internal.messages._fetchSearchMessages, { searchAllMessagesForUserId: args.searchAllMessagesForUserId, threadId: args.threadId, vectorIds, textSearchMessages: textSearchMessages?.filter((m) => !vectorIds.includes(m.embeddingId)), 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) => { 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 = (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 !== 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 = {}; 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 = {}; 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) .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), }); //# sourceMappingURL=messages.js.map