UNPKG

@convex-dev/agent

Version:

A agent component for Convex.

388 lines (370 loc) 10.5 kB
import { omit, pick } from "convex-helpers"; import { v } from "convex/values"; import { type StreamDelta, vStreamDelta, vStreamMessage, } from "../validators.js"; import { api, internal } from "./_generated/api.js"; import type { Id } from "./_generated/dataModel.js"; import { internalMutation, mutation, type MutationCtx, query, action, } from "./_generated/server.js"; import schema from "./schema.js"; import { stream } from "convex-helpers/server/stream"; import { mergedStream } from "convex-helpers/server/stream"; import { paginator } from "convex-helpers/server/pagination"; const SECOND = 1000; const MINUTE = 60 * SECOND; const MAX_DELTAS_PER_REQUEST = 1000; const MAX_DELTAS_PER_STREAM = 100; const TIMEOUT_INTERVAL = 10 * MINUTE; const DELETE_STREAM_DELAY = MINUTE * 5; // 5 minutes const deltaValidator = schema.tables.streamDeltas.validator; export const addDelta = mutation({ args: deltaValidator, returns: v.boolean(), handler: async (ctx, args) => { await ctx.db.insert("streamDeltas", args); await heartbeatStream(ctx, { streamId: args.streamId }); const stream = await ctx.db.get(args.streamId); if (stream?.state.kind !== "streaming") { console.warn(`Stream is not streaming: ${args.streamId}`); return false; } return true; }, }); export const listDeltas = query({ args: { threadId: v.id("threads"), cursors: v.array( v.object({ streamId: v.id("streamingMessages"), cursor: v.number(), }) ), }, returns: v.array(vStreamDelta), handler: async (ctx, args): Promise<StreamDelta[]> => { let totalDeltas = 0; const deltas: StreamDelta[] = []; for (const cursor of args.cursors) { const streamDeltas = await ctx.db .query("streamDeltas") .withIndex("streamId_start_end", (q) => q.eq("streamId", cursor.streamId).gte("start", cursor.cursor) ) .take( Math.min(MAX_DELTAS_PER_STREAM, MAX_DELTAS_PER_REQUEST - totalDeltas) ); totalDeltas += streamDeltas.length; deltas.push( ...streamDeltas.map((d) => pick(d, ["streamId", "start", "end", "parts"]) ) ); if (totalDeltas >= MAX_DELTAS_PER_REQUEST) { break; } } return deltas; }, }); export const create = mutation({ args: omit(schema.tables.streamingMessages.validator.fields, ["state"]), returns: v.id("streamingMessages"), handler: async (ctx, args) => { const state = { kind: "streaming" as const, lastHeartbeat: Date.now(), }; const streamId = await ctx.db.insert("streamingMessages", { ...args, state, }); const timeoutFnId = await ctx.scheduler.runAfter( TIMEOUT_INTERVAL, internal.streams.timeoutStream, { streamId } ); await ctx.db.patch(streamId, { state: { ...state, timeoutFnId } }); return streamId; }, }); export const list = query({ args: { threadId: v.id("threads"), }, returns: v.array(vStreamMessage), handler: async (ctx, args) => { return ctx.db .query("streamingMessages") .withIndex("threadId_state_order_stepOrder", (q) => q.eq("threadId", args.threadId).eq("state.kind", "streaming") ) .order("desc") .take(100) .then((msgs) => msgs.map((m) => ({ streamId: m._id, ...pick(m, [ "order", "stepOrder", "userId", "agentName", "model", "provider", "providerOptions", ]), })) ); }, }); export const finish = mutation({ args: { streamId: v.id("streamingMessages"), finalDelta: v.optional(deltaValidator), }, returns: v.null(), handler: async (ctx, args) => { if (args.finalDelta) { await ctx.db.insert("streamDeltas", args.finalDelta); } const stream = await ctx.db.get(args.streamId); if (!stream) { throw new Error(`Stream not found: ${args.streamId}`); } if (stream.state.kind !== "streaming") { console.warn( `Stream trying to finish but not currently streaming: ${args.streamId}` ); return; } if (stream.state.timeoutFnId) { const timeoutFn = await ctx.db.system.get(stream.state.timeoutFnId); if (timeoutFn?.state.kind === "pending") { await ctx.scheduler.cancel(stream.state.timeoutFnId); } } await ctx.db.patch(args.streamId, { state: { kind: "finished", endedAt: Date.now() }, }); await ctx.scheduler.runAfter( DELETE_STREAM_DELAY, api.streams.deleteStreamAsync, { streamId: args.streamId } ); }, }); async function heartbeatStream( ctx: MutationCtx, args: { streamId: Id<"streamingMessages"> } ) { const stream = await ctx.db.get(args.streamId); if (!stream) { console.warn("Stream not found", args.streamId); return; } if (stream.state.kind !== "streaming") { console.warn("Stream is not streaming", args.streamId); return; } if (Date.now() - stream.state.lastHeartbeat < TIMEOUT_INTERVAL / 4) { // Debounce heartbeating. return; } if (!stream.state.timeoutFnId) { throw new Error("Stream has no timeout function"); } const timeoutFn = await ctx.db.system.get(stream.state.timeoutFnId); if (!timeoutFn) { throw new Error("Timeout function not found"); } if (timeoutFn.state.kind !== "pending") { throw new Error("Timeout function is not pending"); } await ctx.scheduler.cancel(stream.state.timeoutFnId); const timeoutFnId = await ctx.scheduler.runAfter( TIMEOUT_INTERVAL, internal.streams.timeoutStream, { streamId: args.streamId } ); await ctx.db.patch(args.streamId, { state: { kind: "streaming", lastHeartbeat: Date.now(), timeoutFnId, }, }); } export const timeoutStream = internalMutation({ args: { streamId: v.id("streamingMessages") }, returns: v.null(), handler: async (ctx, args) => { const stream = await ctx.db.get(args.streamId); if (!stream) { console.warn("Stream not found", args.streamId); return; } await ctx.db.patch(args.streamId, { state: { kind: "finished", endedAt: Date.now(), }, }); }, }); async function deletePageForStreamId( ctx: MutationCtx, args: { streamId: Id<"streamingMessages">; cursor?: string } ) { const deltas = await paginator(ctx.db, schema) .query("streamDeltas") .withIndex("streamId_start_end", (q) => q.eq("streamId", args.streamId)) .paginate({ numItems: MAX_DELTAS_PER_REQUEST, cursor: args.cursor ?? null, }); await Promise.all(deltas.page.map((d) => ctx.db.delete(d._id))); if (deltas.isDone) { await ctx.db.delete(args.streamId); } return deltas; } export async function deleteStreamsPageForThreadId( ctx: MutationCtx, args: { threadId: Id<"threads">; streamOrder?: number; deltaCursor?: string } ) { const allStreamMessages = schema.tables.streamingMessages.validator.fields.state.members .flatMap((state) => state.fields.kind.value) .map((stateKind) => stream(ctx.db, schema) .query("streamingMessages") .withIndex("threadId_state_order_stepOrder", (q) => q .eq("threadId", args.threadId) .eq("state.kind", stateKind) .gte("order", args.streamOrder ?? 0) ) ); let deltaCursor = args.deltaCursor; const streamMessage = await mergedStream(allStreamMessages, [ "threadId", "state.kind", "order", "stepOrder", ]).first(); if (!streamMessage) { return { isDone: true, streamOrder: undefined, deltaCursor: undefined, }; } const result = await deletePageForStreamId(ctx, { streamId: streamMessage._id, cursor: deltaCursor, }); if (result.isDone) { deltaCursor = undefined; } return { isDone: false, streamOrder: streamMessage.order, deltaCursor, }; } export const deleteStreamsPageForThreadIdMutation = internalMutation({ args: { threadId: v.id("threads"), streamOrder: v.optional(v.number()), deltaCursor: v.optional(v.string()), }, returns: v.object({ isDone: v.boolean(), streamOrder: v.optional(v.number()), deltaCursor: v.optional(v.string()), }), handler: deleteStreamsPageForThreadId, }); export const deleteAllStreamsForThreadIdAsync = mutation({ args: { threadId: v.id("threads"), streamOrder: v.optional(v.number()), deltaCursor: v.optional(v.string()), }, returns: v.object({ isDone: v.boolean(), streamOrder: v.optional(v.number()), deltaCursor: v.optional(v.string()), }), handler: async (ctx, args) => { const result = await deleteStreamsPageForThreadId(ctx, args); if (!result.isDone) { await ctx.scheduler.runAfter( 0, api.streams.deleteAllStreamsForThreadIdAsync, { threadId: args.threadId, streamOrder: result.streamOrder, deltaCursor: result.deltaCursor, } ); } else { await ctx.db.delete(args.threadId); } return result; }, }); export const deleteStreamSync = mutation({ args: { streamId: v.id("streamingMessages") }, returns: v.null(), handler: async (ctx, args) => { let deltas = await deletePageForStreamId(ctx, args); while (!deltas.isDone) { deltas = await deletePageForStreamId(ctx, { ...args, cursor: deltas.continueCursor, }); } }, }); export const deleteStreamAsync = mutation({ args: { streamId: v.id("streamingMessages"), cursor: v.optional(v.string()) }, returns: v.null(), handler: async (ctx, args) => { const result = await deletePageForStreamId(ctx, args); if (!result.isDone) { await ctx.scheduler.runAfter(0, api.streams.deleteStreamAsync, { streamId: args.streamId, cursor: result.continueCursor, }); } }, }); export const deleteAllStreamsForThreadIdSync = action({ args: { threadId: v.id("threads") }, returns: v.null(), handler: async (ctx, args) => { let result = await ctx.runMutation( internal.streams.deleteStreamsPageForThreadIdMutation, args ); while (!result.isDone) { result = await ctx.runMutation( internal.streams.deleteStreamsPageForThreadIdMutation, { ...args, streamOrder: result.streamOrder, deltaCursor: result.deltaCursor, } ); } }, });