@convex-dev/agent
Version:
A agent component for Convex.
388 lines (370 loc) • 10.5 kB
text/typescript
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,
}
);
}
},
});