@convex-dev/agent
Version:
A agent component for Convex.
1,146 lines • 48.7 kB
JavaScript
import { generateObject, generateText, streamObject, streamText } from "ai";
import { assert } from "convex-helpers";
import { internalActionGeneric, internalMutationGeneric, } from "convex/server";
import { v } from "convex/values";
import { validateVectorDimension, } from "../component/vector/tables.js";
import { deserializeMessage, promptOrMessagesToCoreMessages, serializeMessage, serializeNewMessagesInStep, serializeObjectResult, } from "../mapping.js";
import { DEFAULT_MESSAGE_RANGE, DEFAULT_RECENT_MESSAGES, extractText, isTool, } from "../shared.js";
import { vMessageWithMetadata, vSafeObjectArgs, vTextArgs, } from "../validators.js";
import { createTool, wrapTools } from "./createTool.js";
import { DeltaStreamer, mergeTransforms, } from "./streaming.js";
export { storeFile, getFile } from "./files.js";
export { serializeDataOrUrl } from "../mapping.js";
export { vMessageDoc, vThreadDoc } from "../component/schema.js";
export { vAssistantMessage, vContextOptions, vMessage, vPaginationResult, vProviderMetadata, vStorageOptions, vStreamArgs, vSystemMessage, vToolMessage, vUsage, vUserMessage, } from "../validators.js";
export { createTool, extractText, isTool };
export class Agent {
component;
options;
constructor(component, options) {
this.component = component;
this.options = options;
}
async createThread(ctx, args) {
const threadDoc = await ctx.runMutation(this.component.threads.createThread, {
userId: args?.userId,
title: args?.title,
summary: args?.summary,
});
if (!("runAction" in ctx)) {
return { threadId: threadDoc._id };
}
const { thread } = await this.continueThread(ctx, {
threadId: threadDoc._id,
userId: args?.userId,
usageHandler: args?.usageHandler,
tools: args?.tools,
});
return {
threadId: threadDoc._id,
thread,
};
}
/**
* Continues a thread using this agent. Note: threads can be continued
* by different agents. This is a convenience around calling the various
* generate and stream functions with explicit userId and threadId parameters.
* @param ctx The ctx object passed to the action handler
* @param { threadId, userId }: the thread and user to associate the messages with.
* @returns Functions bound to the userId and threadId on a `{thread}` object.
*/
async continueThread(ctx, args) {
return {
thread: {
threadId: args.threadId,
getMetadata: this.getThreadMetadata.bind(this, ctx, {
threadId: args.threadId,
}),
updateMetadata: (patch) => ctx.runMutation(this.component.threads.updateThread, {
threadId: args.threadId,
patch,
}),
generateText: this.generateText.bind(this, ctx, args),
streamText: this.streamText.bind(this, ctx, args),
generateObject: this.generateObject.bind(this, ctx, args),
streamObject: this.streamObject.bind(this, ctx, args),
},
};
}
/**
* This behaves like {@link generateText} from the "ai" package except that
* it add context based on the userId and threadId and saves the input and
* resulting messages to the thread, if specified.
* Use {@link continueThread} to get a version of this function already scoped
* to a thread (and optionally userId).
* @param ctx The context passed from the action function calling this.
* @param { userId, threadId }: The user and thread to associate the message with
* @param args The arguments to the generateText function, along with extra controls
* for the {@link ContextOptions} and {@link StorageOptions}.
* @returns The result of the generateText function.
*/
async generateText(ctx, { userId: argsUserId, threadId, usageHandler, tools: threadTools, }, args, options) {
const context = await this._saveMessagesAndFetchContext(ctx, args, {
userId: argsUserId,
threadId,
...options,
});
const { args: aiArgs, messageId, userId } = context;
const toolCtx = { ...ctx, userId, threadId, messageId, agent: this };
const tools = wrapTools(toolCtx, args.tools ?? threadTools ?? this.options.tools);
const saveOutputMessages = this._shouldSaveOutputMessages(options?.storageOptions);
const trackUsage = usageHandler ?? this.options.usageHandler;
try {
const result = (await generateText({
// Can be overridden
maxSteps: this.options.maxSteps,
...aiArgs,
tools,
onStepFinish: async (step) => {
if (threadId && messageId && saveOutputMessages) {
await this.saveStep(ctx, {
userId,
threadId,
promptMessageId: messageId,
step,
});
}
if (this.options.rawRequestResponseHandler) {
await this.options.rawRequestResponseHandler(ctx, {
userId,
threadId,
agentName: this.options.name,
request: step.request,
response: step.response,
});
}
if (trackUsage && step.usage) {
await trackUsage(ctx, {
userId,
threadId,
agentName: this.options.name,
model: aiArgs.model.modelId,
provider: aiArgs.model.provider,
usage: step.usage,
providerMetadata: step.providerMetadata,
});
}
return args.onStepFinish?.(step);
},
}));
result.messageId = messageId;
return result;
}
catch (error) {
if (threadId && messageId) {
console.error("RollbackMessage", messageId);
await ctx.runMutation(this.component.messages.rollbackMessage, {
messageId,
error: error.message,
});
}
throw error;
}
}
/**
* This behaves like {@link streamText} from the "ai" package except that
* it add context based on the userId and threadId and saves the input and
* resulting messages to the thread, if specified.
* Use {@link continueThread} to get a version of this function already scoped
* to a thread (and optionally userId).
*/
async streamText(ctx, { userId: argsUserId, threadId, usageHandler,
/**
* @deprecated Pass `tools` in the next parameter instead.
* This is only intended to pass through thread-default tools.
*/
tools: threadTools, },
/**
* The arguments to the streamText function, similar to the ai `streamText` function.
*/
args,
/**
* The {@link ContextOptions} and {@link StorageOptions}
* options to use for fetching contextual messages and saving input/output messages.
*/
options) {
const context = await this._saveMessagesAndFetchContext(ctx, args, {
userId: argsUserId,
threadId,
...options,
});
const { args: aiArgs, messageId, order, stepOrder, userId } = context;
const toolCtx = { ...ctx, userId, threadId, messageId, agent: this };
const tools = wrapTools(toolCtx, args.tools ?? threadTools ?? this.options.tools);
const saveOutputMessages = this._shouldSaveOutputMessages(options?.storageOptions);
const trackUsage = usageHandler ?? this.options.usageHandler;
const streamer = threadId && options?.saveStreamDeltas
? new DeltaStreamer(this.component, ctx, options.saveStreamDeltas, {
threadId,
userId,
agentName: this.options.name,
model: aiArgs.model.modelId,
provider: aiArgs.model.provider,
providerOptions: aiArgs.providerOptions,
order,
stepOrder,
abortSignal: aiArgs.abortSignal,
})
: undefined;
const result = streamText({
// Can be overridden
maxSteps: this.options.maxSteps,
...aiArgs,
tools,
abortSignal: streamer?.abortController.signal ?? aiArgs.abortSignal,
experimental_transform: mergeTransforms(options?.saveStreamDeltas, args.experimental_transform),
onChunk: async (event) => {
await streamer?.addParts([event.chunk]);
// console.log("onChunk", chunk);
return args.onChunk?.(event);
},
onError: async (error) => {
console.error("onError", error);
if (threadId && messageId && saveOutputMessages) {
await ctx.runMutation(this.component.messages.rollbackMessage, {
messageId,
error: error.error.message,
});
}
return args.onError?.(error);
},
onStepFinish: async (step) => {
// console.log("onStepFinish", step);
// TODO: compare delta to the output. internally drop the deltas when committing
if (threadId && messageId) {
const saved = await this.saveStep(ctx, {
userId,
threadId,
promptMessageId: messageId,
step,
});
// TODO: figure out pending/not
await streamer?.finish(saved.messages);
}
if (this.options.rawRequestResponseHandler) {
await this.options.rawRequestResponseHandler(ctx, {
userId,
threadId,
agentName: this.options.name,
request: step.request,
response: step.response,
});
}
if (trackUsage && step.usage) {
await trackUsage(ctx, {
userId,
threadId,
agentName: this.options.name,
model: aiArgs.model.modelId,
provider: aiArgs.model.provider,
usage: step.usage,
providerMetadata: step.providerMetadata,
});
}
return args.onStepFinish?.(step);
},
});
result.messageId = messageId;
return result;
}
/**
* This behaves like {@link generateObject} from the "ai" package except that
* it add context based on the userId and threadId and saves the input and
* resulting messages to the thread, if specified.
* Use {@link continueThread} to get a version of this function already scoped
* to a thread (and optionally userId).
*/
async generateObject(ctx, { userId: argsUserId, threadId, usageHandler, },
/**
* The arguments to the generateObject function, similar to the ai.generateObject function.
*/
args,
/**
* The {@link ContextOptions} and {@link StorageOptions}
* options to use for fetching contextual messages and saving input/output messages.
*/
options) {
const context = await this._saveMessagesAndFetchContext(ctx, args, {
userId: argsUserId,
threadId,
...options,
});
const { args: aiArgs, messageId, userId } = context;
const trackUsage = usageHandler ?? this.options.usageHandler;
const saveOutputMessages = this._shouldSaveOutputMessages(options?.storageOptions);
try {
const result = (await generateObject(
// eslint-disable-next-line @typescript-eslint/no-explicit-any
aiArgs));
if (threadId && messageId && saveOutputMessages) {
await this.saveObject(ctx, {
threadId,
promptMessageId: messageId,
result,
userId,
});
}
result.messageId = messageId;
if (this.options.rawRequestResponseHandler) {
await this.options.rawRequestResponseHandler(ctx, {
userId,
threadId,
agentName: this.options.name,
request: result.request,
response: result.response,
});
}
if (trackUsage && result.usage) {
await trackUsage(ctx, {
userId,
threadId,
agentName: this.options.name,
model: aiArgs.model.modelId,
provider: aiArgs.model.provider,
usage: result.usage,
providerMetadata: result.providerMetadata,
});
}
return result;
}
catch (error) {
if (threadId && messageId) {
await ctx.runMutation(this.component.messages.rollbackMessage, {
messageId,
error: error.message,
});
}
throw error;
}
}
/**
* This behaves like `streamObject` from the "ai" package except that
* it add context based on the userId and threadId and saves the input and
* resulting messages to the thread, if specified.
* Use {@link continueThread} to get a version of this function already scoped
* to a thread (and optionally userId).
*/
async streamObject(ctx, { userId: argsUserId, threadId, usageHandler, },
/**
* The arguments to the streamObject function, similar to the ai `streamObject` function.
*/
args,
/**
* The {@link ContextOptions} and {@link StorageOptions}
* options to use for fetching contextual messages and saving input/output messages.
*/
options) {
// TODO: unify all this shared code between all the generate* and stream* functions
const context = await this._saveMessagesAndFetchContext(ctx, args, {
userId: argsUserId,
threadId,
...options,
});
const { args: aiArgs, messageId, userId } = context;
const trackUsage = usageHandler ?? this.options.usageHandler;
const saveOutputMessages = this._shouldSaveOutputMessages(options?.storageOptions);
const stream = streamObject({
// eslint-disable-next-line @typescript-eslint/no-explicit-any
...aiArgs,
onError: async (error) => {
console.error("onError", error);
return args.onError?.(error);
},
onFinish: async (result) => {
if (threadId && messageId && saveOutputMessages) {
await this.saveObject(ctx, {
userId,
threadId,
promptMessageId: messageId,
result: {
object: result.object,
finishReason: "stop",
usage: result.usage,
warnings: result.warnings,
request: await stream.request,
response: result.response,
providerMetadata: result.providerMetadata,
experimental_providerMetadata: result.experimental_providerMetadata,
logprobs: undefined,
toJsonResponse: stream.toTextStreamResponse,
},
});
}
if (trackUsage && result.usage) {
await trackUsage(ctx, {
userId,
threadId,
agentName: this.options.name,
model: aiArgs.model.modelId,
provider: aiArgs.model.provider,
usage: result.usage,
providerMetadata: result.providerMetadata,
});
}
if (this.options.rawRequestResponseHandler) {
await this.options.rawRequestResponseHandler(ctx, {
userId,
threadId,
agentName: this.options.name,
request: await stream.request,
response: result.response,
});
}
// eslint-disable-next-line @typescript-eslint/no-explicit-any
return args.onFinish?.(result);
},
});
stream.messageId = messageId;
return stream;
}
/**
* Save a message to the thread.
* @param ctx A ctx object from a mutation or action.
* @param args The message and what to associate it with (user / thread)
* You can pass extra metadata alongside the message, e.g. associated fileIds.
* @returns The messageId of the saved message.
*/
async saveMessage(ctx, args) {
const { lastMessageId, messages } = await this.saveMessages(ctx, {
threadId: args.threadId,
userId: args.userId,
messages: args.prompt !== undefined
? [{ role: "user", content: args.prompt }]
: [args.message],
metadata: args.metadata ? [args.metadata] : undefined,
skipEmbeddings: args.skipEmbeddings,
});
return { messageId: lastMessageId, message: messages.at(-1) };
}
/**
* Explicitly save messages associated with the thread (& user if provided)
* @param ctx The ctx parameter to a mutation or action.
* @param args The messages and context to save
* @returns
*/
async saveMessages(ctx, args) {
let embeddings;
if (args.skipEmbeddings || !("runAction" in ctx)) {
embeddings = undefined;
if (!args.skipEmbeddings && this.options.textEmbedding) {
console.warn("You're trying to save messages and generate embeddings, but you're in a mutation. " +
"Pass `skipEmbeddings: true` to skip generating embeddings in the mutation and skip this warning. " +
"They will be generated lazily when you generate or stream text / objects. " +
"You can explicitly generate them asynchronously by using the scheduler to run an action later that calls `agent.generateAndSaveEmbeddings`.");
}
}
else {
embeddings = await this.generateEmbeddings(ctx, {
userId: args.userId,
threadId: args.threadId,
}, args.messages);
}
const result = await ctx.runMutation(this.component.messages.addMessages, {
threadId: args.threadId,
userId: args.userId,
agentName: this.options.name,
promptMessageId: args.promptMessageId,
embeddings,
messages: await Promise.all(args.messages.map(async (m, i) => {
const { message, fileIds } = await serializeMessage(ctx, this.component, m);
return {
...args.metadata?.[i],
message,
fileIds,
};
})),
failPendingSteps: args.failPendingSteps ?? false,
pending: args.pending ?? false,
});
return {
lastMessageId: result.messages.at(-1)._id,
messages: result.messages,
};
}
/**
* List messages from a thread.
* @param ctx A ctx object from a query, mutation, or action.
* @param args.threadId The thread to list messages from.
* @param args.paginationOpts Pagination options (e.g. via usePaginatedQuery).
* @param args.excludeToolMessages Whether to exclude tool messages.
* False by default.
* @param args.statuses What statuses to include. All by default.
* @returns The MessageDoc's in a format compatible with usePaginatedQuery.
*/
async listMessages(ctx, args) {
if (args.paginationOpts.numItems === 0) {
return {
page: [],
isDone: true,
continueCursor: args.paginationOpts.cursor ?? "",
};
}
return ctx.runQuery(this.component.messages.listMessagesByThreadId, {
order: "desc",
...args,
});
}
/**
* A function that handles fetching stream deltas, used with the React hooks
* `useThreadMessages` or `useStreamingThreadMessages`.
* @param ctx A ctx object from a query, mutation, or action.
* @param args.threadId The thread to sync streams for.
* @param args.streamArgs The stream arguments with per-stream cursors.
* @returns The deltas for each stream from their existing cursor.
*/
async syncStreams(ctx, args) {
if (!args.streamArgs)
return undefined;
if (args.streamArgs.kind === "list") {
return {
kind: "list",
messages: await ctx.runQuery(this.component.streams.list, {
threadId: args.threadId,
}),
};
}
else {
return {
kind: "deltas",
deltas: await ctx.runQuery(this.component.streams.listDeltas, {
threadId: args.threadId,
cursors: args.streamArgs.cursors,
}),
};
}
}
/**
* Fetch the context messages for a thread.
* @param ctx Either a query, mutation, or action ctx.
* If it is not an action context, you can't do text or
* vector search.
* @param args The associated thread, user, message
* @returns
*/
async fetchContextMessages(ctx, args) {
assert(args.userId || args.threadId, "Specify userId or threadId");
// Fetch the latest messages from the thread
let included;
const opts = this._mergedContextOptions(args.contextOptions);
const contextMessages = [];
if (args.threadId &&
(opts.recentMessages !== 0 || args.upToAndIncludingMessageId)) {
const { page } = await ctx.runQuery(this.component.messages.listMessagesByThreadId, {
threadId: args.threadId,
excludeToolMessages: opts.excludeToolMessages,
paginationOpts: {
numItems: opts.recentMessages ?? DEFAULT_RECENT_MESSAGES,
cursor: null,
},
upToAndIncludingMessageId: args.upToAndIncludingMessageId,
order: "desc",
statuses: ["success"],
});
included = new Set(page.map((m) => m._id));
contextMessages.push(
// Reverse since we fetched in descending order
...page.reverse());
}
if (opts.searchOptions?.textSearch || opts.searchOptions?.vectorSearch) {
const targetMessage = contextMessages.find((m) => m._id === args.upToAndIncludingMessageId)?.message;
const messagesToSearch = targetMessage
? [targetMessage, ...args.messages]
: args.messages;
if (!("runAction" in ctx)) {
throw new Error("searchUserMessages only works in an action");
}
const searchMessages = await ctx.runAction(this.component.messages.searchMessages, {
searchAllMessagesForUserId: opts?.searchOtherThreads
? args.userId ??
(args.threadId &&
(await ctx.runQuery(this.component.threads.getThread, {
threadId: args.threadId,
}))?.userId)
: undefined,
threadId: args.threadId,
beforeMessageId: args.upToAndIncludingMessageId,
...(await this._searchOptionsWithEmbeddingAndDefaults(ctx, { userId: args.userId, threadId: args.threadId }, opts, messagesToSearch)),
});
// TODO: track what messages we used for context
contextMessages.unshift(...searchMessages.filter((m) => !included?.has(m._id)));
}
// Ensure we don't include tool messages without a corresponding tool call
return filterOutOrphanedToolMessages(contextMessages.sort((a, b) =>
// Sort the raw MessageDocs by order and stepOrder
a.order === b.order ? a.stepOrder - b.stepOrder : a.order - b.order));
}
/**
* Get the metadata for a thread.
* @param ctx A ctx object from a query, mutation, or action.
* @param args.threadId The thread to get the metadata for.
* @returns The metadata for the thread.
*/
async getThreadMetadata(ctx, args) {
const thread = await ctx.runQuery(this.component.threads.getThread, {
threadId: args.threadId,
});
if (!thread) {
throw new Error("Thread not found");
}
return thread;
}
/**
* Update the metadata for a thread.
* @param ctx A ctx object from a mutation or action.
* @param args.threadId The thread to update the metadata for.
* @param args.patch The patch to apply to the thread.
* @returns The updated thread metadata.
*/
async updateThreadMetadata(ctx, args) {
const thread = await ctx.runMutation(this.component.threads.updateThread, args);
return thread;
}
/**
* Get the embeddings for a set of messages.
* @param messages The messages to get the embeddings for.
* @returns The embeddings for the messages.
*/
async generateEmbeddings(ctx, { userId, threadId, }, messages) {
if (!this.options.textEmbedding) {
return undefined;
}
let embeddings;
const messageTexts = messages.map((m) => !isTool(m) && extractText(m));
// Find the indexes of the messages that have text.
const textIndexes = messageTexts
.map((t, i) => (t ? i : undefined))
.filter((i) => i !== undefined);
if (textIndexes.length === 0) {
return undefined;
}
// Then embed those messages.
const textEmbeddings = await this.doEmbed(ctx, {
userId,
threadId,
values: messageTexts.filter((t) => !!t),
});
// TODO: record usage of embeddings
// Then assemble the embeddings into a single array with nulls for the messages without text.
const embeddingsOrNull = Array(messages.length).fill(null);
textIndexes.forEach((i, j) => {
embeddingsOrNull[i] = textEmbeddings.embeddings[j];
});
if (textEmbeddings.embeddings.length > 0) {
const dimension = textEmbeddings.embeddings[0].length;
validateVectorDimension(dimension);
embeddings = {
vectors: embeddingsOrNull,
dimension,
model: this.options.textEmbedding.modelId,
};
}
return embeddings;
}
/**
* Generate embeddings for a set of messages, and save them to the database.
* It will not generate or save embeddings for messages that already have an
* embedding.
* @param ctx The ctx parameter to an action.
* @param args The messageIds to generate embeddings for.
*/
async generateAndSaveEmbeddings(ctx, args) {
const messages = (await ctx.runQuery(this.component.messages.getMessagesByIds, {
messageIds: args.messageIds,
})).filter((m) => m !== null);
if (messages.length !== args.messageIds.length) {
throw new Error("Some messages were not found: " +
args.messageIds
.filter((id) => !messages.some((m) => m?._id === id))
.join(", "));
}
if (messages.some((m) => !m.message)) {
throw new Error("Some messages don't have a message: " +
args.messageIds
.map((id, i) => (!messages[i].message ? id : undefined))
.filter((id) => id !== undefined)
.join(", "));
}
const messagesMissingEmbeddings = messages.filter((m) => !m.embeddingId);
if (messagesMissingEmbeddings.length === 0) {
return;
}
const embeddings = await this.generateEmbeddings(ctx, {
userId: messagesMissingEmbeddings[0].userId,
threadId: messagesMissingEmbeddings[0].threadId,
}, messagesMissingEmbeddings.map((m) => m.message));
if (!embeddings) {
if (!this.options.textEmbedding) {
throw new Error("No embeddings were generated for the messages. You must pass a textEmbedding model to the agent constructor.");
}
throw new Error("No embeddings were generated for these messages: " +
messagesMissingEmbeddings.map((m) => m._id).join(", "));
}
await ctx.runMutation(this.component.vector.index.insertBatch, {
vectorDimension: embeddings.dimension,
vectors: messagesMissingEmbeddings
.map((m, i) => ({
messageId: m._id,
model: embeddings.model,
table: "messages",
userId: m.userId,
threadId: m.threadId,
vector: embeddings.vectors[i],
}))
.filter((v) => v.vector !== null),
});
}
/**
* Explicitly save a "step" created by the AI SDK.
* @param ctx The ctx argument to a mutation or action.
* @param args The Step generated by the AI SDK.
*/
async saveStep(ctx, args) {
const messages = await serializeNewMessagesInStep(ctx, this.component, args.step, {
provider: args.provider ?? this.options.chat.provider,
model: args.model ?? this.options.chat.modelId,
});
const embeddings = await this.generateEmbeddings(ctx, { userId: args.userId, threadId: args.threadId }, messages.map((m) => m.message));
const saved = await ctx.runMutation(this.component.messages.addMessages, {
userId: args.userId,
threadId: args.threadId,
agentName: this.options.name,
promptMessageId: args.promptMessageId,
messages,
embeddings,
failPendingSteps: false,
});
return saved;
}
/**
* Manually save the result of a generateObject call to the thread.
* This happens automatically when using {@link generateObject} or {@link streamObject}
* from the `thread` object created by {@link continueThread} or {@link createThread}.
* @param ctx The context passed from the mutation or action function calling this.
* @param args The arguments to the saveObject function.
*/
async saveObject(ctx, args) {
const { messages } = serializeObjectResult(args.result, {
model: this.options.chat.modelId,
provider: this.options.chat.provider,
});
const embeddings = await this.generateEmbeddings(ctx, { userId: args.userId, threadId: args.threadId }, messages.map((m) => m.message));
await ctx.runMutation(this.component.messages.addMessages, {
userId: args.userId,
threadId: args.threadId,
promptMessageId: args.promptMessageId,
failPendingSteps: false,
messages,
embeddings,
agentName: this.options.name,
pending: false,
});
}
/**
* Commit or rollback a message that was pending.
* This is done automatically when saving messages by default.
* If creating pending messages, you can call this when the full "transaction" is done.
* @param ctx The ctx argument to your mutation or action.
* @param args What message to save. Generally the parent message sent into
* the generateText call.
*/
async completeMessage(ctx, args) {
const result = args.result;
if (result.kind === "success") {
await ctx.runMutation(this.component.messages.commitMessage, {
messageId: args.messageId,
});
}
else {
await ctx.runMutation(this.component.messages.rollbackMessage, {
messageId: args.messageId,
error: result.error,
});
}
}
async _saveMessagesAndFetchContext(ctx, args, { userId: argsUserId, threadId, contextOptions, storageOptions, }) {
contextOptions ||= this.options.contextOptions;
storageOptions ||= this.options.storageOptions;
// If only a messageId is provided, this will be empty.
const messages = args.promptMessageId
? []
: promptOrMessagesToCoreMessages(args);
const userId = argsUserId ??
(threadId &&
(await ctx.runQuery(this.component.threads.getThread, { threadId }))
?.userId);
assert(!args.promptMessageId || !(args.prompt || args.messages), "you can't specify a prompt or message if you specify a promptMessageId");
// If only a messageId is provided, this will add that message to the end.
const contextMessages = await this.fetchContextMessages(ctx, {
userId,
threadId,
upToAndIncludingMessageId: args.promptMessageId,
messages,
contextOptions,
});
// Lazily generate embeddings for the prompt message, if it doesn't have
// embeddings yet. This can happen if the message was saved in a mutation
// where the LLM is not available.
if (args.promptMessageId &&
!contextMessages.at(-1)?.embeddingId &&
this.options.textEmbedding) {
await this.generateAndSaveEmbeddings(ctx, {
messageIds: [args.promptMessageId],
});
}
let messageId = args.promptMessageId;
let order = args.promptMessageId
? contextMessages.at(-1)?.order
: undefined;
let stepOrder = args.promptMessageId
? contextMessages.at(-1)?.stepOrder
: undefined;
if (threadId &&
messages.length &&
storageOptions?.saveMessages !== "none" &&
storageOptions?.saveAnyInputMessages !== false) {
const saveAll = storageOptions?.saveMessages === "all";
const coreMessages = saveAll ? messages : messages.slice(-1);
const saved = await this.saveMessages(ctx, {
threadId,
userId,
messages: coreMessages,
metadata: coreMessages.length === 1 ? [{ id: args.id }] : undefined,
pending: true,
failPendingSteps: true,
});
messageId = saved.lastMessageId;
order = saved.messages.at(-1)?.order;
stepOrder = saved.messages.at(-1)?.stepOrder;
}
let processedMessages = [
...contextMessages.map((m) => deserializeMessage(m.message)),
...messages,
];
// Process messages to inline localhost files (if not, file urls pointing to localhost will be sent to LLM providers)
if (process.env.CONVEX_CLOUD_URL?.startsWith("http://127.0.0.1")) {
processedMessages = await this._inlineMessagesFiles(processedMessages);
}
const { prompt: _, model, ...rest } = args;
return {
args: {
...rest,
maxRetries: args.maxRetries ?? this.options.maxRetries,
model: model ?? this.options.chat,
system: args.system ?? this.options.instructions,
messages: processedMessages,
},
userId,
messageId,
order,
stepOrder,
};
}
_shouldSaveOutputMessages(storageOpts) {
const opts = storageOpts ?? this.options.storageOptions;
return opts?.saveOutputMessages !== false && opts?.saveMessages !== "none";
}
_mergedContextOptions(opts) {
const searchOptions = {
...this.options.contextOptions?.searchOptions,
...opts?.searchOptions,
};
return {
...this.options.contextOptions,
...opts,
searchOptions: searchOptions.limit
? searchOptions
: undefined,
};
}
async _searchOptionsWithEmbeddingAndDefaults(ctx, { userId, threadId }, contextOptions, messages) {
assert(contextOptions.searchOptions?.textSearch ||
contextOptions.searchOptions?.vectorSearch, "searchOptions is required");
assert(messages.length > 0, "Core messages cannot be empty");
const text = extractText(messages.at(-1));
const search = {
limit: contextOptions.searchOptions?.limit ?? 10,
messageRange: {
...DEFAULT_MESSAGE_RANGE,
...contextOptions.searchOptions?.messageRange,
},
text: extractText(messages.at(-1)),
};
if (contextOptions.searchOptions?.vectorSearch &&
text &&
this.options.textEmbedding) {
search.vector = (await this.doEmbed(ctx, {
threadId,
userId,
values: [text],
})).embeddings[0];
search.vectorModel = this.options.textEmbedding.modelId;
}
return search;
}
async doEmbed(ctx, options) {
const embedding = this.options.textEmbedding;
assert(embedding, "textEmbedding is required");
const result = await embedding.doEmbed({
values: options.values,
abortSignal: options.abortSignal,
headers: options.headers,
});
if (this.options.usageHandler && result.usage) {
await this.options.usageHandler(ctx, {
userId: options.userId,
threadId: options.threadId,
agentName: this.options.name,
model: embedding.modelId,
provider: embedding.provider,
providerMetadata: result.rawResponse
? { [embedding.provider]: result.rawResponse }
: undefined,
usage: {
promptTokens: result.usage.tokens,
completionTokens: 0,
totalTokens: result.usage.tokens,
},
});
}
return { embeddings: result.embeddings };
}
/**
* Process messages to inline file and image URLs that point to localhost
* by converting them to base64. This solves the problem of LLMs not being
* able to access localhost URLs.
*/
async _inlineMessagesFiles(messages) {
// Process each message to convert localhost URLs to base64
return Promise.all(messages.map(async (message) => {
if (message.role !== "user" ||
typeof message.content === "string" ||
!Array.isArray(message.content)) {
return message;
}
const processedContent = await Promise.all(message.content.map(async (part) => {
if (part.type === "image" && part.image instanceof URL) {
if (this._isLocalhostUrl(part.image)) {
const imageData = await this._downloadFile(part.image);
return {
...part,
image: imageData,
};
}
}
// Handle file parts
if (part.type === "file" && part.data instanceof URL) {
if (this._isLocalhostUrl(part.data)) {
const fileData = await this._downloadFile(part.data);
return {
...part,
data: fileData,
};
}
}
return part;
}));
return {
...message,
content: processedContent,
};
}));
}
/**
* Check if a URL points to localhost
*/
_isLocalhostUrl(url) {
return (url.hostname === "localhost" ||
url.hostname === "127.0.0.1" ||
url.hostname === "::1" ||
url.hostname === "0.0.0.0");
}
/**
* Download a file from a URL
*/
async _downloadFile(url) {
// Fetch the file
const response = await fetch(url);
if (!response.ok) {
throw new Error(`Failed to fetch ${url}: ${response.statusText}`);
}
return await response.arrayBuffer();
}
/**
* WORKFLOW UTILITIES
*/
/**
* Create a mutation that creates a thread so you can call it from a Workflow.
* e.g.
* ```ts
* // in convex/foo.ts
* export const createThread = weatherAgent.createThreadMutation();
*
* const workflow = new WorkflowManager(components.workflow);
* export const myWorkflow = workflow.define({
* args: {},
* handler: async (step) => {
* const { threadId } = await step.runMutation(internal.foo.createThread);
* // use the threadId to generate text, object, etc.
* },
* });
* ```
* @returns A mutation that creates a thread.
*/
createThreadMutation() {
return internalMutationGeneric({
args: {
userId: v.optional(v.string()),
title: v.optional(v.string()),
summary: v.optional(v.string()),
},
handler: async (ctx, args) => {
const { threadId } = await this.createThread(ctx, args);
return { threadId };
},
});
}
/**
* Create an action out of this agent so you can call it from workflows or other actions
* without a wrapping function.
* @param spec Configuration for the agent acting as an action, including
* {@link ContextOptions}, {@link StorageOptions}, and maxSteps.
*/
asTextAction(spec) {
const maxSteps = spec?.maxSteps ?? this.options.maxSteps;
return internalActionGeneric({
args: vTextArgs,
handler: async (ctx, args) => {
const { contextOptions, storageOptions, ...rest } = args;
const stream = args.stream === true ? spec?.stream || true : spec?.stream ?? false;
const targetArgs = { userId: args.userId, threadId: args.threadId };
const llmArgs = { maxSteps, ...rest };
const opts = {
contextOptions: contextOptions ??
spec?.contextOptions ??
this.options.contextOptions,
storageOptions: storageOptions ??
spec?.storageOptions ??
this.options.storageOptions,
saveStreamDeltas: stream,
};
if (stream) {
const result = await this.streamText(ctx, targetArgs, llmArgs, opts);
await result.consumeStream();
return {
text: await result.text,
finishReason: await result.finishReason,
messageId: result.messageId,
};
}
else {
const { text, messageId, finishReason } = await this.generateText(ctx, targetArgs, llmArgs, opts);
return { text, messageId, finishReason };
}
},
});
}
/**
* Create an action that generates an object out of this agent so you can call
* it from workflows or other actions without a wrapping function.
* @param spec Configuration for the agent acting as an action, including
* the normal parameters to {@link generateObject}, plus {@link ContextOptions}
* and maxSteps.
*/
asObjectAction(spec, options) {
const maxSteps = spec?.maxSteps ?? this.options.maxSteps;
return internalActionGeneric({
args: vSafeObjectArgs,
handler: async (ctx, args) => {
const { contextOptions, storageOptions, ...rest } = args;
const value = await this.generateObject(ctx, { userId: args.userId, threadId: args.threadId }, {
...spec,
maxSteps,
...rest,
}, {
contextOptions: contextOptions ??
options?.contextOptions ??
this.options.contextOptions,
storageOptions: storageOptions ??
options?.storageOptions ??
this.options.storageOptions,
});
return { object: value.object };
},
});
}
/**
* Save messages to the thread.
* Useful as a step in Workflows, e.g.
* ```ts
* const saveMessages = agent.asSaveMessagesMutation();
*
* const myWorkflow = workflow.define({
* args: {...},
* handler: async (step, args) => {
* // do things to create (but not save)messages
* const { messageIds } = await step.runMutation(internal.foo.saveMessages, {
* threadId: args.threadId,
* messages: args.messages,
* });
* // ...
* },
* })
* ```
* @returns A mutation that can be used to save messages to the thread.
*/
asSaveMessagesMutation() {
return internalMutationGeneric({
args: {
threadId: v.string(),
userId: v.optional(v.string()),
promptMessageId: v.optional(v.string()),
messages: v.array(vMessageWithMetadata),
pending: v.optional(v.boolean()),
failPendingSteps: v.optional(v.boolean()),
},
handler: async (ctx, args) => {
const { lastMessageId, messages } = await this.saveMessages(ctx, {
...args,
messages: args.messages.map((m) => m.message),
metadata: args.messages.map(({ message: _, ...m }) => m),
});
return {
lastMessageId,
messageIds: messages.map((m) => m._id),
};
},
});
}
}
export function filterOutOrphanedToolMessages(docs) {
const toolCallIds = new Set();
const result = [];
for (const doc of docs) {
if (doc.message?.role === "assistant" &&
Array.isArray(doc.message.content)) {
for (const content of doc.message.content) {
if (content.type === "tool-call") {
toolCallIds.add(content.toolCallId);
}
}
result.push(doc);
}
else if (doc.message?.role === "tool") {
if (doc.message.content.every((c) => toolCallIds.has(c.toolCallId))) {
result.push(doc);
}
else {
console.debug("Filtering out orphaned tool message", doc);
}
}
else {
result.push(doc);
}
}
return result;
}
//# sourceMappingURL=index.js.map