open-swe/apps/web/src/components/thread/messages/ai.tsx
2025-07-14 20:01:27 -07:00

759 lines
24 KiB
TypeScript

import { parsePartialJson } from "@langchain/core/output_parsers";
import {
AIMessage,
Checkpoint,
Message,
ToolMessage,
} from "@langchain/langgraph-sdk";
import { getContentString } from "../utils";
import { BranchSwitcher, CommandBar } from "./shared";
import { MarkdownText } from "../markdown-text";
import {
LoadExternalComponent,
UIMessage,
} from "@langchain/langgraph-sdk/react-ui";
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 {
MarkTaskCompleted,
MarkTaskIncomplete,
} from "@/components/gen-ui/task-review";
import { DiagnoseErrorAction } from "@/components/v2/diagnose-error-action";
import { WriteTechnicalNotes } from "@/components/gen-ui/write-technical-notes";
import { CodeReviewStarted } from "@/components/gen-ui/code-review-started";
import { ToolCall } from "@langchain/core/messages/tool";
import {
createApplyPatchToolFields,
createShellToolFields,
createMarkTaskCompletedToolFields,
createMarkTaskNotCompletedToolFields,
createSearchToolFields,
createOpenPrToolFields,
createInstallDependenciesToolFields,
createTakePlannerNotesFields,
createCodeReviewMarkTaskCompletedFields,
createCodeReviewMarkTaskNotCompleteFields,
createDiagnoseErrorToolFields,
createGetURLContentToolFields,
createWriteTechnicalNotesToolFields,
createConversationHistorySummaryToolFields,
createReviewStartedToolFields,
} 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 reviewStartedTool = createReviewStartedToolFields();
type ReviewStartedToolArgs = z.infer<typeof reviewStartedTool.schema>;
const searchTool = createSearchToolFields(dummyRepo);
type SearchToolArgs = z.infer<typeof searchTool.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>;
const markFinalReviewTaskCompletedTool =
createCodeReviewMarkTaskCompletedFields();
type MarkFinalReviewTaskCompletedToolArgs = z.infer<
typeof markFinalReviewTaskCompletedTool.schema
>;
const markFinalReviewTaskIncompleteTool =
createCodeReviewMarkTaskNotCompleteFields();
type MarkFinalReviewTaskIncompleteToolArgs = z.infer<
typeof markFinalReviewTaskIncompleteTool.schema
>;
const diagnoseErrorTool = createDiagnoseErrorToolFields();
type DiagnoseErrorToolArgs = z.infer<typeof diagnoseErrorTool.schema>;
const getURLContentTool = createGetURLContentToolFields();
type GetURLContentToolArgs = z.infer<typeof getURLContentTool.schema>;
const writeTechnicalNotesTool = createWriteTechnicalNotesToolFields();
type WriteTechnicalNotesToolArgs = z.infer<
typeof writeTechnicalNotesTool.schema
>;
const conversationHistorySummaryTool =
createConversationHistorySummaryToolFields();
type ConversationHistorySummaryToolArgs = z.infer<
typeof conversationHistorySummaryTool.schema
>;
function CustomComponent({
message,
thread,
}: {
message: Message;
thread: ReturnType<typeof useStream>;
}) {
const values = thread.values;
const customComponents =
"ui" in values
? (values.ui as UIMessage[]).filter(
(ui) => ui.metadata?.message_id === message.id,
)
: [];
if (!customComponents?.length) return null;
return (
<Fragment key={message.id}>
{customComponents.map((customComponent) => (
<LoadExternalComponent
key={customComponent.id}
stream={thread}
message={customComponent}
meta={{ ui: customComponent }}
/>
))}
</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 === searchTool.name) {
const args = toolCall.args as SearchToolArgs;
return {
actionType: "search",
status,
success,
query: args.query || "",
match_string: args.match_string || false,
case_sensitive: args.case_sensitive || false,
context_lines: args.context_lines || 0,
max_results: args.max_results || 0,
follow_symlinks: args.follow_symlinks || false,
exclude_files: args.exclude_files || "",
include_files: args.include_files || "",
file_types: args.file_types || [],
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,
};
}
return {
status: "loading",
summaryText: reasoningText,
};
}
export function AssistantMessage({
message,
isLoading,
handleRegenerate,
forceRenderInterrupt = false,
thread,
threadMessages,
}: {
message: Message | undefined;
isLoading: boolean;
handleRegenerate: (parentCheckpoint: Checkpoint | null | undefined) => void;
forceRenderInterrupt?: boolean;
thread: ReturnType<typeof useStream<Record<string, unknown>>>;
threadMessages: Message[];
}) {
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;
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 === searchTool.name ||
tc.name === installDependenciesTool.name ||
tc.name === plannerNotesTool.name ||
tc.name === getURLContentTool.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 markFinalReviewTaskCompletedToolCall = message
? aiToolCalls.find(
(tc) => tc.name === markFinalReviewTaskCompletedTool.name,
)
: undefined;
const markFinalReviewTaskIncompleteToolCall = message
? aiToolCalls.find(
(tc) => tc.name === markFinalReviewTaskIncompleteTool.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;
const reviewStartedToolCall = message
? aiToolCalls.find((tc) => tc.name === reviewStartedTool.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>
);
}
// Check if this is a review started message
if (reviewStartedToolCall && aiToolCalls.length === 1) {
const correspondingToolResult = toolResults.find(
(tr) => tr && tr.tool_call_id === reviewStartedToolCall.id,
);
return (
<div className="flex flex-col gap-4">
<CodeReviewStarted
status={correspondingToolResult ? "done" : "generating"}
/>
</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 task completed review tool call is present, render the task review component
if (markFinalReviewTaskCompletedToolCall) {
const args =
markFinalReviewTaskCompletedToolCall.args as MarkFinalReviewTaskCompletedToolArgs;
const correspondingToolResult = toolResults.find(
(tr) => tr && tr.tool_call_id === markFinalReviewTaskCompletedToolCall.id,
);
const status = correspondingToolResult ? "done" : "generating";
return (
<div className="flex flex-col gap-4">
<MarkTaskCompleted
status={status}
review={args.review}
reasoningText={contentString}
/>
</div>
);
}
// If task incomplete review tool call is present, render the task review component
if (markFinalReviewTaskIncompleteToolCall) {
const args =
markFinalReviewTaskIncompleteToolCall.args as MarkFinalReviewTaskIncompleteToolArgs;
const correspondingToolResult = toolResults.find(
(tr) =>
tr && tr.tool_call_id === markFinalReviewTaskIncompleteToolCall.id,
);
const status = correspondingToolResult ? "done" : "generating";
return (
<div className="flex flex-col gap-4">
<MarkTaskIncomplete
status={status}
review={args.review}
additionalActions={args.additional_actions}
reasoningText={contentString}
/>
</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 isSearchTool = toolCall.name === searchTool.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 (isSearchTool) {
const args = toolCall.args as SearchToolArgs;
return {
actionType: "search",
status: "generating",
query: args?.query || "",
match_string: args?.match_string || false,
case_sensitive: args?.case_sensitive || false,
context_lines: args?.context_lines || 0,
max_results: args?.max_results || 0,
follow_symlinks: args?.follow_symlinks || false,
exclude_files: args?.exclude_files || [],
include_files: args?.include_files || [],
file_types: args?.file_types || [],
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 (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",
);
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">
{isToolResult ? (
<span>
<ToolResult message={message} />
<Interrupt
interruptValue={threadInterrupt?.value}
isLastMessage={isLastMessage}
hasNoAIOrToolMessages={hasNoAIOrToolMessages}
forceRenderInterrupt={forceRenderInterrupt}
thread={thread}
/>
</span>
) : (
<span>
{contentString.length > 0 && (
<div className="py-1">
<MarkdownText>{contentString}</MarkdownText>
</div>
)}
{!hideToolCalls && aiToolCalls.length > 0 && (
<span>
<ToolCalls toolCalls={aiToolCalls} />
</span>
)}
{message && (
<CustomComponent
message={message}
thread={thread}
/>
)}
<Interrupt
interruptValue={threadInterrupt?.value}
isLastMessage={isLastMessage}
hasNoAIOrToolMessages={hasNoAIOrToolMessages}
forceRenderInterrupt={forceRenderInterrupt}
thread={thread}
/>
<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>
)}
</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>
);
}