open-swe/apps/web/src/components/thread/messages/ai.tsx

691 lines
22 KiB
TypeScript
Raw Normal View History

2025-05-28 12:26:47 -07:00
import { parsePartialJson } from "@langchain/core/output_parsers";
import {
AIMessage,
Checkpoint,
Message,
ToolMessage,
} from "@langchain/langgraph-sdk";
2025-05-28 12:26:47 -07:00
import { getContentString } from "../utils";
import { BranchSwitcher, CommandBar } from "./shared";
import { MarkdownText } from "../markdown-text";
import {
LoadExternalComponent,
UIMessage,
} from "@langchain/langgraph-sdk/react-ui";
2025-05-28 12:26:47 -07:00
import { cn } from "@/lib/utils";
import { ToolCalls, ToolResult } from "./tool-calls";
import { MessageContentComplex } from "@langchain/core/messages";
import { Fragment } from "react/jsx-runtime";
import { useQueryState, parseAsBoolean } from "nuqs";
import { Interrupt } from "./interrupt";
import { ActionStep, ActionItemProps } from "@/components/gen-ui/action-step";
import { TaskSummary } from "@/components/gen-ui/task-summary";
import { PullRequestOpened } from "@/components/gen-ui/pull-request-opened";
import { DiagnoseErrorAction } from "@/components/v2/diagnose-error-action";
import { WriteTechnicalNotes } from "@/components/gen-ui/write-technical-notes";
import { ToolCall } from "@langchain/core/messages/tool";
import {
createApplyPatchToolFields,
createShellToolFields,
createMarkTaskCompletedToolFields,
createMarkTaskNotCompletedToolFields,
createRgToolFields,
createOpenPrToolFields,
createInstallDependenciesToolFields,
createTakePlannerNotesFields,
createDiagnoseErrorToolFields,
createGetURLContentToolFields,
createFindInstancesOfToolFields,
createWriteTechnicalNotesToolFields,
createConversationHistorySummaryToolFields,
} from "@open-swe/shared/open-swe/tools";
import { z } from "zod";
import { isAIMessageSDK, isToolMessageSDK } from "@/lib/langchain-messages";
import { useStream } from "@langchain/langgraph-sdk/react";
import { ConversationHistorySummary } from "@/components/gen-ui/conversation-summary";
// Used only for Zod type inference.
const dummyRepo = { owner: "dummy", repo: "dummy" };
const shellTool = createShellToolFields(dummyRepo);
type ShellToolArgs = z.infer<typeof shellTool.schema>;
const applyPatchTool = createApplyPatchToolFields(dummyRepo);
type ApplyPatchToolArgs = z.infer<typeof applyPatchTool.schema>;
const markTaskCompletedTool = createMarkTaskCompletedToolFields();
type MarkTaskCompletedToolArgs = z.infer<typeof markTaskCompletedTool.schema>;
const markTaskNotCompletedTool = createMarkTaskNotCompletedToolFields();
type MarkTaskNotCompletedToolArgs = z.infer<
typeof markTaskNotCompletedTool.schema
>;
const rgTool = createRgToolFields(dummyRepo);
type RgToolArgs = z.infer<typeof rgTool.schema>;
const openPrTool = createOpenPrToolFields();
type OpenPrToolArgs = z.infer<typeof openPrTool.schema>;
const installDependenciesTool = createInstallDependenciesToolFields(dummyRepo);
type InstallDependenciesToolArgs = z.infer<
typeof installDependenciesTool.schema
>;
const plannerNotesTool = createTakePlannerNotesFields();
type PlannerNotesToolArgs = z.infer<typeof plannerNotesTool.schema>;
2025-05-28 12:26:47 -07:00
const diagnoseErrorTool = createDiagnoseErrorToolFields();
type DiagnoseErrorToolArgs = z.infer<typeof diagnoseErrorTool.schema>;
const getURLContentTool = createGetURLContentToolFields();
type GetURLContentToolArgs = z.infer<typeof getURLContentTool.schema>;
const findInstancesOfTool = createFindInstancesOfToolFields(dummyRepo);
type FindInstancesOfToolArgs = z.infer<typeof findInstancesOfTool.schema>;
const writeTechnicalNotesTool = createWriteTechnicalNotesToolFields();
type WriteTechnicalNotesToolArgs = z.infer<
typeof writeTechnicalNotesTool.schema
>;
const conversationHistorySummaryTool =
createConversationHistorySummaryToolFields();
type ConversationHistorySummaryToolArgs = z.infer<
typeof conversationHistorySummaryTool.schema
>;
2025-05-28 12:26:47 -07:00
function CustomComponent({
message,
thread,
}: {
message: Message;
thread: ReturnType<typeof useStream>;
2025-05-28 12:26:47 -07:00
}) {
const values = thread.values;
const customComponents =
"ui" in values
? (values.ui as UIMessage[]).filter(
(ui) => ui.metadata?.message_id === message.id,
)
: [];
2025-05-28 12:26:47 -07:00
if (!customComponents?.length) return null;
return (
<Fragment key={message.id}>
{customComponents.map((customComponent) => (
<LoadExternalComponent
key={customComponent.id}
stream={thread}
message={customComponent}
meta={{ ui: customComponent }}
2025-05-28 12:26:47 -07:00
/>
))}
</Fragment>
);
}
function parseAnthropicStreamedToolCalls(
content: MessageContentComplex[],
): AIMessage["tool_calls"] {
const toolCallContents = content.filter((c) => c.type === "tool_use" && c.id);
return toolCallContents.map((tc) => {
const toolCall = tc as Record<string, any>;
let json: Record<string, any> = {};
if (toolCall?.input) {
try {
json = parsePartialJson(toolCall.input) ?? {};
} catch {
// Pass
}
}
return {
name: toolCall.name ?? "",
id: toolCall.id ?? "",
args: json,
type: "tool_call",
};
});
}
export function mapToolMessageToActionStepProps(
message: ToolMessage,
threadMessages: Message[],
): ActionItemProps {
const toolCall: ToolCall | undefined = threadMessages
.filter(isAIMessageSDK)
.flatMap((m) => m.tool_calls ?? [])
.find((tc) => tc.id === message.tool_call_id);
const aiMessage = threadMessages
.filter(isAIMessageSDK)
.find((m) => m.tool_calls?.some((tc) => tc.id === message.tool_call_id));
const reasoningText = aiMessage
? getContentString(aiMessage.content)
: undefined;
const status: ActionItemProps["status"] = "done";
const success = message.status === "success";
if (toolCall?.name === shellTool.name) {
const args = toolCall.args as ShellToolArgs;
return {
actionType: shellTool.name as "shell",
status,
success,
command: args.command || [],
workdir: args.workdir,
output: getContentString(message.content),
reasoningText,
};
} else if (toolCall?.name === applyPatchTool.name) {
const args = toolCall.args as ApplyPatchToolArgs;
return {
actionType: "apply-patch",
status,
success,
file_path: args.file_path || "",
diff: args.diff,
reasoningText,
errorMessage: !success ? getContentString(message.content) : undefined,
};
} else if (toolCall?.name === rgTool.name) {
const args = toolCall.args as RgToolArgs;
return {
actionType: "rg",
status,
success,
pattern: args.pattern || "",
paths: args.paths || [],
output: getContentString(message.content),
reasoningText,
};
} else if (toolCall?.name === installDependenciesTool.name) {
const args = toolCall.args as InstallDependenciesToolArgs;
return {
actionType: "install_dependencies",
status,
success,
command: args.command || "",
workdir: args.workdir || "",
output: getContentString(message.content),
reasoningText,
};
} else if (toolCall?.name === plannerNotesTool.name) {
const args = toolCall.args as PlannerNotesToolArgs;
return {
actionType: "planner_notes",
status,
success,
notes: args.notes || [],
reasoningText,
};
} else if (toolCall?.name === getURLContentTool.name) {
const args = toolCall.args as GetURLContentToolArgs;
return {
actionType: "get_url_content",
status,
success,
url: args.url || "",
output: getContentString(message.content),
reasoningText,
};
} else if (toolCall?.name === findInstancesOfTool.name) {
const args = toolCall.args as FindInstancesOfToolArgs;
// case_sensitive and match_word both default to true.
const caseSensitive =
args.case_sensitive === undefined ? true : args.case_sensitive;
const matchWord = args.match_word === undefined ? true : args.match_word;
return {
actionType: "find_instances_of",
status,
success,
query: args.query || "",
case_sensitive: caseSensitive,
match_word: matchWord,
include_files: args.include_files,
exclude_files: args.exclude_files,
output: getContentString(message.content),
reasoningText,
};
}
return {
status: "loading",
summaryText: reasoningText,
};
}
2025-05-28 12:26:47 -07:00
export function AssistantMessage({
message,
isLoading,
handleRegenerate,
forceRenderInterrupt = false,
thread,
threadMessages,
2025-05-28 12:26:47 -07:00
}: {
message: Message | undefined;
isLoading: boolean;
handleRegenerate: (parentCheckpoint: Checkpoint | null | undefined) => void;
forceRenderInterrupt?: boolean;
thread: ReturnType<typeof useStream<Record<string, unknown>>>;
threadMessages: Message[];
2025-05-28 12:26:47 -07:00
}) {
const content = message?.content ?? [];
const contentString = getContentString(content);
const [hideToolCalls] = useQueryState(
"hideToolCalls",
parseAsBoolean.withDefault(false),
);
const messages = threadMessages;
const idx = message ? messages.findIndex((m) => m.id === message.id) : -1;
2025-05-28 12:26:47 -07:00
const meta = message ? thread.getMessagesMetadata(message) : undefined;
const threadInterrupt = thread.interrupt;
const parentCheckpoint = meta?.firstSeenState?.parent_checkpoint;
const anthropicStreamedToolCalls = Array.isArray(content)
? parseAnthropicStreamedToolCalls(content)
: undefined;
const aiToolCalls: ToolCall[] = (() => {
if (message && isAIMessageSDK(message)) {
return message.tool_calls || [];
}
if (anthropicStreamedToolCalls?.length) {
return anthropicStreamedToolCalls;
}
return [];
})();
const toolResults = aiToolCalls
.map((toolCall) => {
const matchingToolMessage = messages.find(
(m) => isToolMessageSDK(m) && m.tool_call_id === toolCall.id,
);
return matchingToolMessage as ToolMessage | undefined;
})
.filter((m): m is ToolMessage => !!m);
const actionableToolCalls = message
? aiToolCalls.filter(
(tc) =>
tc.name === shellTool.name ||
tc.name === applyPatchTool.name ||
tc.name === rgTool.name ||
tc.name === installDependenciesTool.name ||
tc.name === plannerNotesTool.name ||
tc.name === getURLContentTool.name ||
tc.name === findInstancesOfTool.name,
)
: [];
const markTaskCompletedToolCall = message
? aiToolCalls.find((tc) => tc.name === markTaskCompletedTool.name)
: undefined;
const markTaskNotCompletedToolCall = message
? aiToolCalls.find((tc) => tc.name === markTaskNotCompletedTool.name)
: undefined;
const openPrToolCall = message
? aiToolCalls.find((tc) => tc.name === openPrTool.name)
: undefined;
const diagnoseErrorToolCall = message
? aiToolCalls.find((tc) => tc.name === diagnoseErrorTool.name)
: undefined;
const writeTechnicalNotesToolCall = message
? aiToolCalls.find((tc) => tc.name === writeTechnicalNotesTool.name)
: undefined;
const conversationHistorySummaryToolCall = message
? aiToolCalls.find((tc) => tc.name === conversationHistorySummaryTool.name)
: undefined;
// Check if this is a conversation history summary message
if (conversationHistorySummaryToolCall && aiToolCalls.length === 1) {
const args =
conversationHistorySummaryToolCall.args as ConversationHistorySummaryToolArgs;
return (
<div className="flex flex-col gap-4">
<ConversationHistorySummary
summary={args.conversation_history_summary}
/>
</div>
);
}
// We can be sure that if either task status tool call is present, it will be the
// only tool call/result we need to render for this message.
if (markTaskCompletedToolCall || markTaskNotCompletedToolCall) {
const toolCall = markTaskCompletedToolCall || markTaskNotCompletedToolCall;
const completed = !!markTaskCompletedToolCall;
const correspondingToolResult = toolResults.find(
(tr) => tr && tr.tool_call_id === toolCall!.id,
);
const status = correspondingToolResult ? "done" : "generating";
// Get the appropriate summary text based on which tool was called
const summaryText = markTaskCompletedToolCall
? (markTaskCompletedToolCall.args as MarkTaskCompletedToolArgs)
.completed_task_summary
: (markTaskNotCompletedToolCall!.args as MarkTaskNotCompletedToolArgs)
.reasoning;
return (
<div className="flex flex-col gap-4">
<TaskSummary
status={status}
completed={completed}
summaryText={summaryText}
/>
</div>
);
}
if (diagnoseErrorToolCall) {
const correspondingToolResult = toolResults.find(
(tr) => tr && tr.tool_call_id === diagnoseErrorToolCall.id,
);
const args = diagnoseErrorToolCall.args as DiagnoseErrorToolArgs;
const reasoningText = getContentString(content);
return (
<div className="flex flex-col gap-4">
<DiagnoseErrorAction
status={correspondingToolResult ? "done" : "generating"}
diagnosis={args.diagnosis}
reasoningText={reasoningText}
/>
</div>
);
}
if (writeTechnicalNotesToolCall) {
const correspondingToolResult = toolResults.find(
(tr) => tr && tr.tool_call_id === writeTechnicalNotesToolCall.id,
);
const args =
writeTechnicalNotesToolCall.args as WriteTechnicalNotesToolArgs;
const reasoningText = getContentString(content);
return (
<div className="flex flex-col gap-4">
<WriteTechnicalNotes
status={correspondingToolResult ? "done" : "generating"}
notes={args.notes}
reasoningText={reasoningText}
/>
</div>
);
}
if (openPrToolCall) {
let branch: string | undefined;
let targetBranch: string | undefined = "main";
if (message && isAIMessageSDK(message)) {
branch = message.additional_kwargs?.branch as string | undefined;
targetBranch =
(message.additional_kwargs?.targetBranch as string | undefined) ||
"main";
}
const args = openPrToolCall.args as OpenPrToolArgs;
const correspondingToolResult = toolResults.find(
(tr) => tr && tr.tool_call_id === openPrToolCall.id,
);
const status = correspondingToolResult ? "done" : "generating";
// Extract PR URL from the tool message content
// Format: "Created pull request: https://github.com/owner/repo/pull/123"
let prUrl: string | undefined = undefined;
if (correspondingToolResult) {
const content = getContentString(correspondingToolResult.content);
if (content.includes("Created pull request: ")) {
prUrl = content.split("Created pull request: ")[1].trim();
}
}
// Extract PR number from URL if available
let prNumber: number | undefined = undefined;
if (prUrl) {
const match = prUrl.match(/\/pull\/(\d+)/);
if (match && match[1]) {
prNumber = parseInt(match[1], 10);
}
}
return (
<div className="flex flex-col gap-4">
<PullRequestOpened
status={status}
title={args.title}
description={args.body}
url={prUrl}
prNumber={prNumber}
branch={branch}
targetBranch={targetBranch}
/>
</div>
);
}
if (actionableToolCalls.length > 0) {
const actionItems = actionableToolCalls.map((toolCall): ActionItemProps => {
const correspondingToolResult = toolResults.find(
(tr) => tr && tr.tool_call_id === toolCall.id,
);
const isShellTool = toolCall.name === shellTool.name;
const isRgTool = toolCall.name === rgTool.name;
const isInstallDependenciesTool =
toolCall.name === installDependenciesTool.name;
if (correspondingToolResult) {
// If we have a tool result, map it to action props
return mapToolMessageToActionStepProps(
correspondingToolResult,
threadMessages,
);
} else if (isRgTool) {
const args = toolCall.args as RgToolArgs;
return {
actionType: "rg",
status: "generating",
pattern: args?.pattern || "",
paths: args?.paths || [],
output: "",
} as ActionItemProps;
} else if (isInstallDependenciesTool) {
const args = toolCall.args as InstallDependenciesToolArgs;
return {
actionType: "install_dependencies",
status: "generating",
command: args?.command || "",
workdir: args?.workdir || "",
output: "",
} as ActionItemProps;
} else if (toolCall.name === plannerNotesTool.name) {
const args = toolCall.args as PlannerNotesToolArgs;
return {
actionType: "planner_notes",
status: "generating",
notes: args?.notes || [],
} as ActionItemProps;
} else if (toolCall.name === getURLContentTool.name) {
const args = toolCall.args as GetURLContentToolArgs;
return {
actionType: "get_url_content",
status: "generating",
url: args?.url || "",
output: "",
} as ActionItemProps;
} else if (toolCall.name === findInstancesOfTool.name) {
const args = toolCall.args as FindInstancesOfToolArgs;
// case_sensitive and match_word both default to true.
const caseSensitive =
args.case_sensitive === undefined ? true : args.case_sensitive;
const matchWord =
args.match_word === undefined ? true : args.match_word;
return {
actionType: "find_instances_of",
status: "generating",
query: args?.query || "",
case_sensitive: caseSensitive,
match_word: matchWord,
include_files: args?.include_files,
exclude_files: args?.exclude_files,
output: "",
} as ActionItemProps;
} else {
if (isShellTool) {
const args = toolCall.args as ShellToolArgs;
return {
actionType: "shell",
status: "generating",
command: args?.command || [],
workdir: args?.workdir,
timeout: args?.timeout,
} as ActionItemProps;
} else {
// Must be apply_patch tool
const patchArgs = toolCall.args as ApplyPatchToolArgs;
return {
actionType: "apply-patch",
status: "generating",
file_path: patchArgs?.file_path || "",
diff: patchArgs?.diff || "",
} as ActionItemProps;
}
}
});
return (
<div className="flex flex-col gap-4">
<ActionStep
actions={actionItems.filter(
(item): item is ActionItemProps => item !== undefined,
)}
reasoningText={contentString}
/>
</div>
);
}
if (message?.type === "tool" && idx > 0) {
const isPreviousToolCall = messages.slice(0, idx).some((prevMessage) => {
if (isAIMessageSDK(prevMessage) && prevMessage.tool_calls) {
return prevMessage.tool_calls.some(
(tc) => tc.id === (message as ToolMessage).tool_call_id,
);
}
if (Array.isArray(prevMessage.content)) {
const toolCalls = parseAnthropicStreamedToolCalls(
prevMessage.content as MessageContentComplex[],
);
return toolCalls?.some(
(tc) => tc.id === (message as ToolMessage).tool_call_id,
);
}
return false;
});
if (isPreviousToolCall) {
return null;
}
}
const isLastMessage =
threadMessages[threadMessages.length - 1].id === message?.id;
const hasNoAIOrToolMessages = !threadMessages.find(
(m) => m.type === "ai" || m.type === "tool",
);
2025-05-28 12:26:47 -07:00
const isToolResult = message?.type === "tool";
if (isToolResult && hideToolCalls) {
return null;
}
return (
<div className="group mr-auto flex w-full items-start gap-2">
<div className="flex w-full flex-col gap-2">
2025-05-28 12:26:47 -07:00
{isToolResult ? (
<span>
2025-05-28 12:26:47 -07:00
<ToolResult message={message} />
<Interrupt
interruptValue={threadInterrupt?.value}
isLastMessage={isLastMessage}
hasNoAIOrToolMessages={hasNoAIOrToolMessages}
forceRenderInterrupt={forceRenderInterrupt}
thread={thread}
2025-05-28 12:26:47 -07:00
/>
</span>
2025-05-28 12:26:47 -07:00
) : (
<span>
2025-05-28 12:26:47 -07:00
{contentString.length > 0 && (
<div className="py-1">
<MarkdownText>{contentString}</MarkdownText>
</div>
)}
{!hideToolCalls && aiToolCalls.length > 0 && (
<span>
<ToolCalls toolCalls={aiToolCalls} />
</span>
2025-05-28 12:26:47 -07:00
)}
{message && (
<CustomComponent
message={message}
thread={thread}
/>
)}
<Interrupt
interruptValue={threadInterrupt?.value}
isLastMessage={isLastMessage}
hasNoAIOrToolMessages={hasNoAIOrToolMessages}
forceRenderInterrupt={forceRenderInterrupt}
thread={thread}
2025-05-28 12:26:47 -07:00
/>
<div
className={cn(
"mr-auto flex items-center gap-2 transition-opacity",
"opacity-0 group-focus-within:opacity-100 group-hover:opacity-100",
)}
>
<BranchSwitcher
branch={meta?.branch}
branchOptions={meta?.branchOptions}
onSelect={(branch) => thread.setBranch(branch)}
isLoading={isLoading}
/>
<CommandBar
content={contentString}
isLoading={isLoading}
isAiMessage={true}
handleRegenerate={() => handleRegenerate(parentCheckpoint)}
/>
</div>
</span>
2025-05-28 12:26:47 -07:00
)}
</div>
</div>
);
}
export function AssistantMessageLoading() {
return (
<div className="mr-auto flex items-start gap-2">
<div className="bg-muted flex h-8 items-center gap-1 rounded-2xl px-4 py-2">
<div className="bg-foreground/50 h-1.5 w-1.5 animate-[pulse_1.5s_ease-in-out_infinite] rounded-full"></div>
<div className="bg-foreground/50 h-1.5 w-1.5 animate-[pulse_1.5s_ease-in-out_0.5s_infinite] rounded-full"></div>
<div className="bg-foreground/50 h-1.5 w-1.5 animate-[pulse_1.5s_ease-in-out_1s_infinite] rounded-full"></div>
</div>
</div>
);
}