@convex-dev/agent
Version:
A agent component for Convex.
342 lines (333 loc) • 9.35 kB
text/typescript
import type { TextPart, ToolCallPart, ToolResultPart } from "ai";
import type { MessageDoc } from "../client/index.js";
import type {
Message,
StreamDelta,
StreamMessage,
TextStreamPart,
} from "../validators.js";
import type { UIMessage } from "./toUIMessages.js";
import { toUIMessages } from "./toUIMessages.js";
export { toUIMessages, type UIMessage };
export function mergeDeltas(
threadId: string,
streamMessages: StreamMessage[],
existingStreams: Array<{
streamId: string;
cursor: number;
messages: MessageDoc[];
}>,
allDeltas: StreamDelta[]
): [
MessageDoc[],
Array<{ streamId: string; cursor: number; messages: MessageDoc[] }>,
boolean,
] {
const newStreams: Array<{
streamId: string;
cursor: number;
messages: MessageDoc[];
}> = [];
// Seed the existing chunks
let changed = false;
for (const streamMessage of streamMessages) {
const deltas = allDeltas.filter(
(d) => d.streamId === streamMessage.streamId
);
const existing = existingStreams.find(
(s) => s.streamId === streamMessage.streamId
);
const [newStream, messageChanged] = applyDeltasToStreamMessage(
threadId,
streamMessage,
existing,
deltas
);
newStreams.push(newStream);
if (messageChanged) changed = true;
}
for (const { streamId } of existingStreams) {
if (!newStreams.find((s) => s.streamId === streamId)) {
// There's a stream that's no longer active.
changed = true;
}
}
const messages = newStreams
.map((s) => s.messages)
.flat()
.sort((a, b) => a.order - b.order || a.stepOrder - b.stepOrder);
return [messages, newStreams, changed];
}
// exported for testing
export function applyDeltasToStreamMessage(
threadId: string,
streamMessage: StreamMessage,
existing:
| { streamId: string; cursor: number; messages: MessageDoc[] }
| undefined,
deltas: StreamDelta[]
): [{ streamId: string; cursor: number; messages: MessageDoc[] }, boolean] {
let changed = false;
let cursor = existing?.cursor ?? 0;
let parts: TextStreamPart[] = [];
for (const delta of deltas.sort((a, b) => a.start - b.start)) {
if (delta.parts.length === 0) {
console.warn(`Got delta with no parts: ${JSON.stringify(delta)}`);
continue;
}
if (cursor !== delta.start) {
if (cursor >= delta.end) {
console.debug(
`Got duplicate delta for stream ${delta.streamId} at ${delta.start}`
);
continue;
} else if (cursor < delta.start) {
console.warn(
`Got delta for stream ${delta.streamId} that has a gap ${cursor} -> ${delta.start}`
);
continue;
} else {
throw new Error(
`Got unexpected delta for stream ${delta.streamId}: delta: ${delta.start} -> ${delta.end} existing cursor: ${cursor}`
);
}
}
changed = true;
cursor = delta.end;
parts.push(...delta.parts);
}
if (!changed) {
return [
existing ?? { streamId: streamMessage.streamId, cursor, messages: [] },
false,
];
}
const existingMessages = existing?.messages ?? [];
let currentMessage: MessageDoc;
if (existingMessages.length > 0) {
// replace the last message with a new one
const lastMessage = existingMessages.at(-1)!;
currentMessage = {
...lastMessage,
message: cloneMessageAndContent(lastMessage.message),
};
} else {
const newMessage = createStreamingMessage(
threadId,
streamMessage,
parts[0]!,
existingMessages.length
);
parts = parts.slice(1);
currentMessage = newMessage;
}
const newStream = {
streamId: streamMessage.streamId,
cursor,
messages: [...existingMessages.slice(0, -1), currentMessage],
};
let lastContent = getLastContent(currentMessage);
for (const part of parts) {
let contentToAdd:
| TextPart
| ToolCallPart
| { type: "reasoning"; text: string }
| ToolResultPart
| undefined;
const isToolRole = part.type === "source" || part.type === "tool-result";
if (isToolRole !== (currentMessage.message!.role === "tool")) {
currentMessage = createStreamingMessage(
threadId,
streamMessage,
part,
newStream.messages.length
);
lastContent = getLastContent(currentMessage);
newStream.messages.push(currentMessage);
continue;
}
switch (part.type) {
case "text-delta":
currentMessage.text += part.textDelta;
if (lastContent?.type === "text") {
lastContent.text += part.textDelta;
} else {
contentToAdd = {
type: "text",
text: part.textDelta,
};
}
break;
case "tool-call-streaming-start":
currentMessage.tool = true;
contentToAdd = {
type: "tool-call",
toolCallId: part.toolCallId,
toolName: part.toolName,
args: "",
};
break;
case "tool-call-delta":
{
currentMessage.tool = true;
if (lastContent?.type !== "tool-call") {
throw new Error("Expected last content to be a tool call");
}
if (typeof lastContent.args !== "string") {
throw new Error("Expected args to be a string");
}
lastContent.args += part.argsTextDelta;
}
break;
case "tool-call":
currentMessage.tool = true;
contentToAdd = part;
break;
case "reasoning":
currentMessage.reasoning += part.textDelta;
if (lastContent?.type === "reasoning") {
lastContent.text += part.textDelta;
} else {
contentToAdd = {
type: "reasoning",
text: part.textDelta,
};
}
break;
case "source":
if (!currentMessage.sources) {
currentMessage.sources = [];
}
currentMessage.sources.push(part.source);
break;
case "tool-result":
contentToAdd = part;
break;
default:
console.warn(`Received unexpected part: ${JSON.stringify(part)}`);
break;
}
if (contentToAdd) {
if (!currentMessage.message!.content) {
currentMessage.message!.content = [];
}
if (!Array.isArray(currentMessage.message?.content)) {
throw new Error("Expected message content to be an array");
}
// eslint-disable-next-line @typescript-eslint/no-explicit-any
currentMessage.message.content.push(contentToAdd as any);
lastContent = contentToAdd;
}
}
return [newStream, true];
}
function cloneMessageAndContent(
message: Message | undefined
): Message | undefined {
return (
message &&
({
...message,
content: Array.isArray(message.content)
? [...message.content]
: message.content,
} as typeof message)
);
}
function getLastContent(message: MessageDoc) {
if (Array.isArray(message.message?.content)) {
return message.message.content.at(-1);
}
return undefined;
}
export function createStreamingMessage(
threadId: string,
message: StreamMessage,
part: TextStreamPart,
index: number
): MessageDoc {
const { streamId, ...rest } = message;
const metadata: MessageDoc = {
_id: `${streamId}-${index}`,
_creationTime: Date.now(),
status: "pending",
threadId,
tool: false,
...rest,
};
switch (part.type) {
case "text-delta":
return {
...metadata,
message: {
role: "assistant",
content: [{ type: "text", text: part.textDelta }],
},
text: part.textDelta,
};
case "tool-call-streaming-start":
return {
...metadata,
tool: true,
message: {
role: "assistant",
content: [
{
type: "tool-call",
toolName: part.toolName,
toolCallId: part.toolCallId,
args: "", // when it's a string, it's a partial call
},
],
},
};
case "reasoning":
return {
...metadata,
message: {
role: "assistant",
content: [{ type: "reasoning", text: part.textDelta }],
},
reasoning: part.textDelta,
};
case "source":
console.warn("Received source part first??");
return {
...metadata,
tool: true,
message: { role: "tool", content: [] },
sources: [part.source],
};
case "tool-call":
return {
...metadata,
tool: true,
message: { role: "assistant", content: [part] },
};
case "tool-call-delta":
console.warn("Received tool call delta part first??");
return {
...metadata,
tool: true,
message: {
role: "assistant",
content: [
{
type: "tool-call",
toolCallId: part.toolCallId,
toolName: part.toolName,
args: part.argsTextDelta,
},
],
},
};
case "tool-result":
return {
...metadata,
tool: true,
message: { role: "tool", content: [part] },
};
default:
throw new Error(`Unexpected part type: ${JSON.stringify(part)}`);
}
}