@convex-dev/agent
Version:
A agent component for Convex.
531 lines (515 loc) • 16.6 kB
text/typescript
import {
readUIMessageStream,
type DynamicToolUIPart,
type ProviderMetadata,
type ReasoningUIPart,
type TextUIPart,
type ToolUIPart,
type UIMessageChunk,
} from "ai";
import { assert } from "convex-helpers";
import { type UIMessage } from "./UIMessages.js";
import { joinText, sorted } from "./shared.js";
import {
type MessageStatus,
type StreamDelta,
type StreamMessage,
} from "./validators.js";
export function blankUIMessage<METADATA = unknown>(
streamMessage: StreamMessage & { metadata?: METADATA },
threadId: string,
): UIMessage<METADATA> {
return {
id: `stream:${streamMessage.streamId}`,
key: `${threadId}-${streamMessage.order}-${streamMessage.stepOrder}`,
order: streamMessage.order,
stepOrder: streamMessage.stepOrder,
status: statusFromStreamStatus(streamMessage.status),
agentName: streamMessage.agentName,
text: "",
_creationTime: Date.now(),
role: "assistant",
parts: [],
...(streamMessage.metadata ? { metadata: streamMessage.metadata } : {}),
};
}
export function statusFromStreamStatus(
status: StreamMessage["status"],
): MessageStatus | "streaming" {
switch (status) {
case "streaming":
return "streaming";
case "finished":
return "success";
case "aborted":
return "failed";
default:
return "pending";
}
}
export async function updateFromUIMessageChunks(
uiMessage: UIMessage,
parts: UIMessageChunk[],
) {
if (parts.length === 0) {
return uiMessage;
}
const partsStream = new ReadableStream<UIMessageChunk>({
start(controller) {
for (const part of parts) {
controller.enqueue(part);
}
controller.close();
},
});
let failed = false;
let suppressError = false;
const messageStream = readUIMessageStream({
message: uiMessage,
stream: partsStream,
onError: (e) => {
const errorMessage = e instanceof Error ? e.message : String(e);
if (errorMessage.toLowerCase().includes("no tool invocation found")) {
suppressError = true;
return;
}
failed = true;
console.error("Error in stream", e);
},
terminateOnError: true,
});
let message = uiMessage;
try {
for await (const messagePart of messageStream) {
assert(
messagePart.id === message.id,
`Expecting to only make one UIMessage in a stream`,
);
message = messagePart;
}
} catch (e) {
if (!suppressError) {
throw e;
}
}
if (failed) {
message.status = "failed";
}
message.text = joinText(message.parts);
return message;
}
type ToolPart = ToolUIPart | DynamicToolUIPart;
function transitionToolPart<S extends ToolPart["state"]>(
part: ToolPart,
updates: { state: S } & Partial<Extract<ToolPart, { state: S }>>,
): void {
Object.assign(part, updates);
}
export type IncrementalStreamState = {
// chunk id -> index of the streaming text part in message.parts
activeText: Record<string, number>;
// chunk id -> index of the streaming reasoning part in message.parts
activeReasoning: Record<string, number>;
// toolCallId -> raw accumulated input JSON text (kept separate from the
// parsed `input` so partial JSON can be repair-parsed each batch)
toolInputText: Record<string, string>;
};
export function emptyIncrementalStreamState(): IncrementalStreamState {
return { activeText: {}, activeReasoning: {}, toolInputText: {} };
}
/**
* Apply a batch of new UIMessageChunks to an existing UIMessage without
* replaying prior chunks. `prev` carries the ephemeral stream state that the
* UIMessage itself can't hold (which text/reasoning parts are still streaming,
* and the raw accumulated tool input text). Parts are append-only, so part
* indices stay stable across the structuredClone between batches. Behavior
* mirrors the AI SDK's processUIMessageStream.
*/
export function applyUIMessageChunksIncremental(
uiMessage: UIMessage,
newParts: UIMessageChunk[],
prev: IncrementalStreamState,
): { message: UIMessage; streamState: IncrementalStreamState } {
const message: UIMessage = structuredClone(uiMessage);
const activeText: Record<string, number> = { ...prev.activeText };
const activeReasoning: Record<string, number> = { ...prev.activeReasoning };
const toolInputText: Record<string, string> = { ...prev.toolInputText };
const touchedTools = new Set<string>();
const toolIndexById = new Map<string, number>();
message.parts.forEach((p, i) => {
if (
"toolCallId" in p &&
(p.type.startsWith("tool-") || p.type === "dynamic-tool")
) {
toolIndexById.set((p as ToolPart).toolCallId, i);
}
});
const toolPartAt = (toolCallId: string): ToolPart | undefined => {
const idx = toolIndexById.get(toolCallId);
return idx === undefined ? undefined : (message.parts[idx] as ToolPart);
};
const mergeMetadata = (metadata: unknown) => {
if (metadata == null) {
return;
}
message.metadata = {
...(message.metadata as Record<string, unknown> | undefined),
...(metadata as Record<string, unknown>),
} as typeof message.metadata;
};
for (const part of newParts) {
switch (part.type) {
case "text-start": {
const newPart: TextUIPart = {
type: "text",
text: "",
state: "streaming",
providerMetadata: part.providerMetadata,
};
message.parts.push(newPart);
activeText[part.id] = message.parts.length - 1;
break;
}
case "text-delta": {
const idx = activeText[part.id];
if (idx !== undefined) {
const textPart = message.parts[idx] as TextUIPart;
textPart.text += part.delta;
textPart.providerMetadata = mergeProviderMetadata(
textPart.providerMetadata,
part.providerMetadata,
);
}
break;
}
case "text-end": {
const idx = activeText[part.id];
if (idx !== undefined) {
const textPart = message.parts[idx] as TextUIPart;
textPart.state = "done";
textPart.providerMetadata = mergeProviderMetadata(
textPart.providerMetadata,
part.providerMetadata,
);
delete activeText[part.id];
}
break;
}
case "reasoning-start": {
const newPart: ReasoningUIPart = {
type: "reasoning",
text: "",
state: "streaming",
providerMetadata: part.providerMetadata,
};
message.parts.push(newPart);
activeReasoning[part.id] = message.parts.length - 1;
break;
}
case "reasoning-delta": {
const idx = activeReasoning[part.id];
if (idx !== undefined) {
const reasoningPart = message.parts[idx] as ReasoningUIPart;
reasoningPart.text += part.delta;
reasoningPart.providerMetadata = mergeProviderMetadata(
reasoningPart.providerMetadata,
part.providerMetadata,
);
}
break;
}
case "reasoning-end": {
const idx = activeReasoning[part.id];
if (idx !== undefined) {
const reasoningPart = message.parts[idx] as ReasoningUIPart;
reasoningPart.state = "done";
reasoningPart.providerMetadata = mergeProviderMetadata(
reasoningPart.providerMetadata,
part.providerMetadata,
);
delete activeReasoning[part.id];
}
break;
}
case "tool-input-start": {
const newToolPart: ToolUIPart | DynamicToolUIPart = part.dynamic
? ({
type: "dynamic-tool",
toolCallId: part.toolCallId,
toolName: part.toolName,
state: "input-streaming",
input: undefined,
} satisfies DynamicToolUIPart)
: ({
type: `tool-${part.toolName}`,
toolCallId: part.toolCallId,
state: "input-streaming",
input: undefined,
providerExecuted: part.providerExecuted,
} satisfies ToolUIPart);
message.parts.push(newToolPart);
toolIndexById.set(part.toolCallId, message.parts.length - 1);
toolInputText[part.toolCallId] = "";
break;
}
case "tool-input-delta": {
if (toolIndexById.has(part.toolCallId)) {
toolInputText[part.toolCallId] =
(toolInputText[part.toolCallId] ?? "") + part.inputTextDelta;
touchedTools.add(part.toolCallId);
} else {
console.warn(
`tool-input-delta for unknown toolCallId ${part.toolCallId}`,
);
}
break;
}
case "tool-input-available": {
const toolPart = toolPartAt(part.toolCallId);
if (toolPart) {
transitionToolPart(toolPart, {
state: "input-available",
input: part.input,
callProviderMetadata: mergeProviderMetadata(
(toolPart as { callProviderMetadata?: ProviderMetadata })
.callProviderMetadata,
part.providerMetadata,
),
});
}
touchedTools.delete(part.toolCallId);
// The raw JSON buffer is no longer needed; drop it so it doesn't get
// carried through every later batch on the hot path.
delete toolInputText[part.toolCallId];
break;
}
case "tool-input-error": {
const toolPart = toolPartAt(part.toolCallId);
if (toolPart) {
transitionToolPart(toolPart, {
state: "output-error",
errorText: part.errorText,
providerExecuted: part.providerExecuted,
...(toolPart.type === "dynamic-tool"
? { input: part.input }
: { input: undefined, rawInput: part.input }),
callProviderMetadata: mergeProviderMetadata(
(toolPart as { callProviderMetadata?: ProviderMetadata })
.callProviderMetadata,
part.providerMetadata,
),
});
}
touchedTools.delete(part.toolCallId);
delete toolInputText[part.toolCallId];
break;
}
case "tool-output-available": {
const toolPart = toolPartAt(part.toolCallId);
if (toolPart) {
transitionToolPart(toolPart, {
state: "output-available",
output: part.output,
preliminary: part.preliminary,
providerExecuted: part.providerExecuted,
});
}
break;
}
case "tool-output-error": {
const toolPart = toolPartAt(part.toolCallId);
if (toolPart) {
transitionToolPart(toolPart, {
state: "output-error",
errorText: part.errorText,
providerExecuted: part.providerExecuted,
});
}
break;
}
case "tool-output-denied": {
const toolPart = toolPartAt(part.toolCallId);
if (toolPart) {
transitionToolPart(toolPart, { state: "output-denied" });
}
break;
}
case "tool-approval-request": {
const toolPart = toolPartAt(part.toolCallId);
if (toolPart) {
transitionToolPart(toolPart, {
state: "approval-requested",
approval: { id: part.approvalId },
});
}
break;
}
case "source-url":
message.parts.push({
type: "source-url",
url: part.url,
sourceId: part.sourceId,
title: part.title,
providerMetadata: part.providerMetadata,
});
break;
case "source-document":
message.parts.push({
type: "source-document",
mediaType: part.mediaType,
sourceId: part.sourceId,
title: part.title,
filename: part.filename,
providerMetadata: part.providerMetadata,
});
break;
case "file":
message.parts.push({
type: "file",
mediaType: part.mediaType,
url: part.url,
});
break;
case "start-step":
message.parts.push({ type: "step-start" });
break;
case "finish-step":
// Match the SDK: a new step starts fresh streaming parts; the prior
// parts keep their state rather than being forced to "done".
for (const id of Object.keys(activeText)) delete activeText[id];
for (const id of Object.keys(activeReasoning))
delete activeReasoning[id];
break;
case "start":
case "finish":
case "message-metadata":
mergeMetadata(part.messageMetadata);
break;
case "abort":
case "error":
// The stream-level status (statusFromStreamStatus) is authoritative and
// is applied by the caller; nothing to mutate on the message here.
break;
default: {
if (typeof part.type === "string" && part.type.startsWith("data-")) {
const dataPart = part as Extract<
UIMessageChunk,
{ type: `data-${string}` }
>;
const existingIdx =
dataPart.id != null
? message.parts.findIndex(
(p) =>
p.type === dataPart.type &&
(p as { id?: string }).id === dataPart.id,
)
: -1;
if (existingIdx >= 0) {
(message.parts[existingIdx] as { data?: unknown }).data =
dataPart.data;
} else {
message.parts.push(
dataPart as unknown as UIMessage["parts"][number],
);
}
} else {
console.warn(
`applyUIMessageChunksIncremental: unhandled chunk type ${String(part.type)}`,
);
}
break;
}
}
}
for (const toolCallId of touchedTools) {
const toolPart = toolPartAt(toolCallId);
if (toolPart && toolPart.state === "input-streaming") {
try {
toolPart.input = JSON.parse(toolInputText[toolCallId] ?? "");
} catch {
// partial JSON — leave input unset until complete
}
}
}
message.text = joinText(message.parts);
return {
message,
streamState: { activeText, activeReasoning, toolInputText },
};
}
export async function deriveUIMessagesFromDeltas(
threadId: string,
streamMessages: StreamMessage[],
allDeltas: StreamDelta[],
): Promise<UIMessage[]> {
const messages: UIMessage[] = [];
for (const streamMessage of streamMessages) {
if (streamMessage.format !== "UIMessageChunk") {
throw new Error(
`deriveUIMessagesFromDeltas: unsupported stream format "${streamMessage.format ?? "text"}" for stream ${streamMessage.streamId}`,
);
}
const { parts } = getParts<UIMessageChunk>(
allDeltas.filter((d) => d.streamId === streamMessage.streamId),
0,
);
const uiMessage = await updateFromUIMessageChunks(
blankUIMessage(streamMessage, threadId),
parts,
);
messages.push(uiMessage);
}
return sorted(messages);
}
export function getParts<T extends StreamDelta["parts"][number]>(
deltas: StreamDelta[],
fromCursor?: number,
): { parts: T[]; cursor: number } {
const parts: T[] = [];
let cursor = fromCursor ?? 0;
for (const delta of deltas.sort((a, b) => a.start - b.start)) {
if (delta.parts.length === 0) {
console.debug(`Got delta with no parts: ${JSON.stringify(delta)}`);
continue;
}
if (cursor !== delta.start) {
if (cursor >= delta.end) {
continue;
} else if (cursor < delta.start) {
console.warn(
`Got delta for stream ${delta.streamId} that has a gap ${cursor} -> ${delta.start}`,
);
break;
} else {
throw new Error(
`Got unexpected delta for stream ${delta.streamId}: delta: ${delta.start} -> ${delta.end} existing cursor: ${cursor}`,
);
}
}
parts.push(...delta.parts);
cursor = delta.end;
}
return { parts, cursor };
}
function mergeProviderMetadata(
existing: ProviderMetadata | undefined,
part: ProviderMetadata | undefined,
): ProviderMetadata | undefined {
if (!existing && !part) {
return undefined;
}
if (!existing) {
return part;
}
if (!part) {
return existing;
}
const merged: ProviderMetadata = existing;
for (const [provider, metadata] of Object.entries(part)) {
merged[provider] = {
...merged[provider],
...metadata,
};
}
return merged;
}