feat: Add auto accept mode (#346)

* feat: Add auto accept mode

* cr
This commit is contained in:
Brace Sproul 2025-07-03 16:01:52 -04:00 • committed by GitHub
parent d7b48c1179
commit 4ea21855f7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 270 additions and 150 deletions

View file

@ -10,6 +10,7 @@ import {
GITHUB_TOKEN_COOKIE, GITHUB_TOKEN_COOKIE,
GITHUB_USER_ID_HEADER, GITHUB_USER_ID_HEADER,
GITHUB_USER_LOGIN_HEADER, GITHUB_USER_LOGIN_HEADER,
MANAGER_GRAPH_ID,
} from "@open-swe/shared/constants"; } from "@open-swe/shared/constants";
import { createLangGraphClient } from "../../../utils/langgraph-client.js"; import { createLangGraphClient } from "../../../utils/langgraph-client.js";
import { createIssue } from "../../../utils/github/api.js"; import { createIssue } from "../../../utils/github/api.js";
@ -91,7 +92,7 @@ ${ISSUE_CONTENT_CLOSE_TAG}`,
messages: inputMessages, messages: inputMessages,
branchName: state.branchName ?? getBranchName(config), branchName: state.branchName ?? getBranchName(config),
}; };
await langGraphClient.runs.create(newManagerThreadId, "manager", { await langGraphClient.runs.create(newManagerThreadId, MANAGER_GRAPH_ID, {
input: {}, input: {},
command: { command: {
update: commandUpdate, update: commandUpdate,

View file

@ -10,9 +10,11 @@ import {
GITHUB_TOKEN_COOKIE, GITHUB_TOKEN_COOKIE,
GITHUB_USER_ID_HEADER, GITHUB_USER_ID_HEADER,
GITHUB_USER_LOGIN_HEADER, GITHUB_USER_LOGIN_HEADER,
PLANNER_GRAPH_ID,
} from "@open-swe/shared/constants"; } from "@open-swe/shared/constants";
import { createLogger, LogLevel } from "../../../utils/logger.js"; import { createLogger, LogLevel } from "../../../utils/logger.js";
import { getBranchName } from "../../../utils/github/git.js"; import { getBranchName } from "../../../utils/github/git.js";
import { PlannerGraphUpdate } from "@open-swe/shared/open-swe/planner/types";
const logger = createLogger(LogLevel.INFO, "StartPlanner"); const logger = createLogger(LogLevel.INFO, "StartPlanner");
@ -38,23 +40,29 @@ export async function startPlanner(
const plannerThreadId = state.plannerSession?.threadId ?? uuidv4(); const plannerThreadId = state.plannerSession?.threadId ?? uuidv4();
try { try {
const run = await langGraphClient.runs.create(plannerThreadId, "planner", { const runInput: PlannerGraphUpdate = {
input: { // github issue ID & target repo so the planning agent can fetch the user's request, and clone the repo.
// github issue ID & target repo so the planning agent can fetch the user's request, and clone the repo. githubIssueId: state.githubIssueId,
githubIssueId: state.githubIssueId, targetRepository: state.targetRepository,
targetRepository: state.targetRepository, // Include the existing task plan, so the agent can use it as context when generating followup tasks.
// Include the existing task plan, so the agent can use it as context when generating followup tasks. taskPlan: state.taskPlan,
taskPlan: state.taskPlan, branchName: state.branchName ?? getBranchName(config),
branchName: state.branchName ?? getBranchName(config), autoAcceptPlan: state.autoAcceptPlan,
};
const run = await langGraphClient.runs.create(
plannerThreadId,
PLANNER_GRAPH_ID,
{
input: runInput,
config: {
recursion_limit: 400,
},
ifNotExists: "create",
multitaskStrategy: "enqueue",
streamResumable: true,
streamMode: ["values", "messages", "custom"],
}, },
config: { );
recursion_limit: 400,
},
ifNotExists: "create",
multitaskStrategy: "enqueue",
streamResumable: true,
streamMode: ["values", "messages", "custom"],
});
return { return {
plannerSession: { plannerSession: {

View file

@ -1,6 +1,10 @@
import { v4 as uuidv4 } from "uuid"; import { v4 as uuidv4 } from "uuid";
import { Command, END, interrupt } from "@langchain/langgraph"; import { Command, END, interrupt } from "@langchain/langgraph";
import { GraphUpdate, GraphConfig } from "@open-swe/shared/open-swe/types"; import {
GraphUpdate,
GraphConfig,
TaskPlan,
} from "@open-swe/shared/open-swe/types";
import { import {
ActionRequest, ActionRequest,
HumanInterrupt, HumanInterrupt,
@ -16,6 +20,7 @@ import {
GITHUB_USER_LOGIN_HEADER, GITHUB_USER_LOGIN_HEADER,
PLAN_INTERRUPT_ACTION_TITLE, PLAN_INTERRUPT_ACTION_TITLE,
PLAN_INTERRUPT_DELIMITER, PLAN_INTERRUPT_DELIMITER,
PROGRAMMER_GRAPH_ID,
} from "@open-swe/shared/constants"; } from "@open-swe/shared/constants";
import { import {
PlannerGraphState, PlannerGraphState,
@ -23,6 +28,63 @@ import {
} from "@open-swe/shared/open-swe/planner/types"; } from "@open-swe/shared/open-swe/planner/types";
import { createLangGraphClient } from "../../../utils/langgraph-client.js"; import { createLangGraphClient } from "../../../utils/langgraph-client.js";
import { addTaskPlanToIssue } from "../../../utils/github/issue-task.js"; import { addTaskPlanToIssue } from "../../../utils/github/issue-task.js";
import { createLogger, LogLevel } from "../../../utils/logger.js";
const logger = createLogger(LogLevel.INFO, "ProposedPlan");
async function startProgrammerRun(input: {
runInput: Exclude<GraphUpdate, "taskPlan"> & { taskPlan: TaskPlan };
state: PlannerGraphState;
config: GraphConfig;
}) {
const { runInput, state, config } = input;
const langGraphClient = createLangGraphClient({
defaultHeaders: {
[GITHUB_TOKEN_COOKIE]: config.configurable?.[GITHUB_TOKEN_COOKIE] ?? "",
[GITHUB_INSTALLATION_TOKEN_COOKIE]:
config.configurable?.[GITHUB_INSTALLATION_TOKEN_COOKIE] ?? "",
[GITHUB_USER_ID_HEADER]:
config.configurable?.[GITHUB_USER_ID_HEADER] ?? "",
[GITHUB_USER_LOGIN_HEADER]:
config.configurable?.[GITHUB_USER_LOGIN_HEADER] ?? "",
},
});
const programmerThreadId = uuidv4();
// Restart the sandbox.
runInput.sandboxSessionId = (await startSandbox(state.sandboxSessionId)).id;
const run = await langGraphClient.runs.create(
programmerThreadId,
PROGRAMMER_GRAPH_ID,
{
input: runInput,
config: {
recursion_limit: 400,
},
ifNotExists: "create",
streamResumable: true,
streamMode: ["values", "messages", "custom"],
},
);
await addTaskPlanToIssue(
{
githubIssueId: state.githubIssueId,
targetRepository: state.targetRepository,
},
config,
runInput.taskPlan,
);
return {
programmerSession: {
threadId: programmerThreadId,
runId: run.run_id,
},
sandboxSessionId: runInput.sandboxSessionId,
taskPlan: runInput.taskPlan,
};
}
export async function interruptProposedPlan( export async function interruptProposedPlan(
state: PlannerGraphState, state: PlannerGraphState,
@ -33,6 +95,37 @@ export async function interruptProposedPlan(
throw new Error("No proposed plan found."); throw new Error("No proposed plan found.");
} }
const userRequest = getUserRequest(state.messages);
const runInput: GraphUpdate = {
contextGatheringNotes: state.contextGatheringNotes,
branchName: state.branchName,
targetRepository: state.targetRepository,
githubIssueId: state.githubIssueId,
};
if (state.autoAcceptPlan) {
logger.info("Auto accepting plan.");
const planItems = proposedPlan.map((p, index) => ({
index,
plan: p,
completed: false,
}));
runInput.taskPlan = createNewTask(
userRequest,
state.proposedPlanTitle,
planItems,
{ existingTaskPlan: state.taskPlan },
);
return await startProgrammerRun({
runInput: runInput as Exclude<GraphUpdate, "taskPlan"> & {
taskPlan: TaskPlan;
},
state,
config,
});
}
const interruptRes = interrupt<HumanInterrupt, HumanResponse[]>({ const interruptRes = interrupt<HumanInterrupt, HumanResponse[]>({
action_request: { action_request: {
action: PLAN_INTERRUPT_ACTION_TITLE, action: PLAN_INTERRUPT_ACTION_TITLE,
@ -67,30 +160,6 @@ export async function interruptProposedPlan(
}); });
} }
const langGraphClient = createLangGraphClient({
defaultHeaders: {
[GITHUB_TOKEN_COOKIE]: config.configurable?.[GITHUB_TOKEN_COOKIE] ?? "",
[GITHUB_INSTALLATION_TOKEN_COOKIE]:
config.configurable?.[GITHUB_INSTALLATION_TOKEN_COOKIE] ?? "",
[GITHUB_USER_ID_HEADER]:
config.configurable?.[GITHUB_USER_ID_HEADER] ?? "",
[GITHUB_USER_LOGIN_HEADER]:
config.configurable?.[GITHUB_USER_LOGIN_HEADER] ?? "",
},
});
const userRequest = getUserRequest(state.messages);
const runInput: GraphUpdate = {
contextGatheringNotes: state.contextGatheringNotes,
branchName: state.branchName,
targetRepository: state.targetRepository,
githubIssueId: state.githubIssueId,
};
// TODO: UPDATE ISSUE WITH PROGRAMMER THREAD ID.
// TODO: UPDATE ISSUE WITH TASK PLAN
const programmerThreadId = uuidv4();
if (interruptRes.type === "accept") { if (interruptRes.type === "accept") {
const planItems = proposedPlan.map((p, index) => ({ const planItems = proposedPlan.map((p, index) => ({
index, index,
@ -125,38 +194,11 @@ export async function interruptProposedPlan(
throw new Error("Unknown interrupt type." + interruptRes.type); throw new Error("Unknown interrupt type." + interruptRes.type);
} }
// Restart the sandbox. return await startProgrammerRun({
runInput.sandboxSessionId = (await startSandbox(state.sandboxSessionId)).id; runInput: runInput as Exclude<GraphUpdate, "taskPlan"> & {
taskPlan: TaskPlan;
const run = await langGraphClient.runs.create(
programmerThreadId,
"programmer",
{
input: runInput,
config: {
recursion_limit: 400,
},
ifNotExists: "create",
streamResumable: true,
streamMode: ["values", "messages", "custom"],
},
);
await addTaskPlanToIssue(
{
githubIssueId: state.githubIssueId,
targetRepository: state.targetRepository,
}, },
state,
config, config,
runInput.taskPlan, });
);
return {
programmerSession: {
threadId: programmerThreadId,
runId: run.run_id,
},
sandboxSessionId: runInput.sandboxSessionId,
taskPlan: runInput.taskPlan,
};
} }

View file

@ -13,7 +13,11 @@ import {
} from "@open-swe/shared/constants"; } from "@open-swe/shared/constants";
import { encryptGitHubToken } from "@open-swe/shared/crypto"; import { encryptGitHubToken } from "@open-swe/shared/crypto";
import { HumanMessage } from "@langchain/core/messages"; import { HumanMessage } from "@langchain/core/messages";
import { getOpenSWELabel } from "../../utils/github/label.js"; import {
getOpenSWEAutoAcceptLabel,
getOpenSWELabel,
} from "../../utils/github/label.js";
import { ManagerGraphUpdate } from "@open-swe/shared/open-swe/manager/types";
const logger = createLogger(LogLevel.INFO, "GitHubIssueWebhook"); const logger = createLogger(LogLevel.INFO, "GitHubIssueWebhook");
@ -72,11 +76,21 @@ webhooks.on("issues.labeled", async ({ payload }) => {
"GITHUB_TOKEN_ENCRYPTION_KEY environment variable is required", "GITHUB_TOKEN_ENCRYPTION_KEY environment variable is required",
); );
} }
if (payload.label?.name !== getOpenSWELabel()) { const validOpenSWELabels = [getOpenSWELabel(), getOpenSWEAutoAcceptLabel()];
if (
!payload.label?.name ||
!validOpenSWELabels.some((l) => l === payload.label?.name)
) {
return; return;
} }
const isAutoAcceptLabel = payload.label.name === getOpenSWEAutoAcceptLabel();
logger.info(`'open-swe' label added to issue #${payload.issue.number}`); logger.info(
`'${payload.label.name}' label added to issue #${payload.issue.number}`,
{
isAutoAcceptLabel,
},
);
try { try {
// Get installation ID from the webhook payload // Get installation ID from the webhook payload
@ -113,24 +127,26 @@ webhooks.on("issues.labeled", async ({ payload }) => {
}); });
const threadId = uuidv4(); const threadId = uuidv4();
const run = await langGraphClient.runs.create(threadId, MANAGER_GRAPH_ID, { const runInput: ManagerGraphUpdate = {
input: { messages: [
messages: [ new HumanMessage({
new HumanMessage({ id: uuidv4(),
id: uuidv4(), content: `**${issueData.issueTitle}**\n\n${issueData.issueBody}`,
content: `**${issueData.issueTitle}**\n\n${issueData.issueBody}`, additional_kwargs: {
additional_kwargs: { isOriginalIssue: true,
isOriginalIssue: true, githubIssueId: issueData.issueNumber,
githubIssueId: issueData.issueNumber, },
}, }),
}), ],
], githubIssueId: issueData.issueNumber,
githubIssueId: issueData.issueNumber, targetRepository: {
targetRepository: { owner: issueData.owner,
owner: issueData.owner, repo: issueData.repo,
repo: issueData.repo,
},
}, },
autoAcceptPlan: isAutoAcceptLabel,
};
const run = await langGraphClient.runs.create(threadId, MANAGER_GRAPH_ID, {
input: runInput,
config: { config: {
recursion_limit: 400, recursion_limit: 400,
}, },
@ -140,13 +156,14 @@ webhooks.on("issues.labeled", async ({ payload }) => {
}); });
logger.info("Created new run from GitHub issue.", { logger.info("Created new run from GitHub issue.", {
thread_id: threadId, threadId,
run_id: run.run_id, runId: run.run_id,
issue_number: issueData.issueNumber, issueNumber: issueData.issueNumber,
owner: issueData.owner, owner: issueData.owner,
repo: issueData.repo, repo: issueData.repo,
user_id: issueData.userId, userId: issueData.userId,
user_login: issueData.userLogin, userLogin: issueData.userLogin,
autoAcceptPlan: isAutoAcceptLabel,
}); });
logger.info("Creating comment..."); logger.info("Creating comment...");

View file

@ -1,6 +1,17 @@
/** /**
* @returns "open-swe" or "open-swe-dev" based on the NODE_ENV. * @returns "open-swe" or "open-swe-dev" based on the NODE_ENV.
*/ */
export function getOpenSWELabel() { export function getOpenSWELabel(): "open-swe" | "open-swe-dev" {
return process.env.NODE_ENV === "production" ? "open-swe" : "open-swe-dev"; return process.env.NODE_ENV === "production" ? "open-swe" : "open-swe-dev";
} }
/**
* @returns "open-swe-auto" or "open-swe-auto-dev" based on the NODE_ENV.
*/
export function getOpenSWEAutoAcceptLabel():
| "open-swe-auto"
| "open-swe-auto-dev" {
return process.env.NODE_ENV === "production"
? "open-swe-auto"
: "open-swe-auto-dev";
}

View file

@ -6,14 +6,13 @@ export function RepositoryBranchSelectors() {
const [threadId] = useQueryState("threadId"); const [threadId] = useQueryState("threadId");
const chatStarted = !!threadId; const chatStarted = !!threadId;
const defaultButtonStyles = const defaultButtonStyles =
"bg-inherit border-none text-gray-500 hover:text-black dark:hover:text-gray-300 text-xs p-0 h-fit hover:bg-inherit"; "bg-inherit border-none text-gray-500 hover:text-black dark:hover:text-gray-300 text-xs p-0 px-0 py-0 !p-0 !px-0 !py-0 h-fit hover:bg-inherit shadow-none";
const defaultStylesChatStarted = const defaultStylesChatStarted =
"hover:bg-inherit cursor-default hover:cursor-default hover:text-black dark:hover:text-gray-300 hover:border-gray-300 hover:ring-inherit"; "hover:bg-inherit cursor-default hover:cursor-default hover:text-black dark:hover:text-gray-300 hover:border-gray-300 hover:ring-inherit shadow-none p-0 px-0 py-0 !p-0 !px-0 !py-0";
return ( return (
<div className="flex items-center gap-2"> <div className="flex items-center gap-1 rounded-md border border-gray-200 p-1">
<div className="flex items-center gap-0"> <div className="flex items-center gap-0">
<span className="-mr-2 text-gray-500">(</span>
<RepositorySelector <RepositorySelector
chatStarted={chatStarted} chatStarted={chatStarted}
buttonClassName={ buttonClassName={
@ -21,10 +20,9 @@ export function RepositoryBranchSelectors() {
(chatStarted ? " " + defaultStylesChatStarted : "") (chatStarted ? " " + defaultStylesChatStarted : "")
} }
/> />
<span className="-ml-2 text-gray-500">)</span>
</div> </div>
<span className="text-muted-foreground/70">:</span>
<div className="flex items-center gap-0"> <div className="flex items-center gap-0">
<span className="-mr-2 text-gray-500">(</span>
<BranchSelector <BranchSelector
chatStarted={chatStarted} chatStarted={chatStarted}
buttonClassName={ buttonClassName={
@ -32,7 +30,6 @@ export function RepositoryBranchSelectors() {
(chatStarted ? " " + defaultStylesChatStarted : "") (chatStarted ? " " + defaultStylesChatStarted : "")
} }
/> />
<span className="-ml-2 text-gray-500">)</span>
</div> </div>
</div> </div>
); );

View file

@ -30,7 +30,8 @@ export type TooltipIconButtonProps = ButtonProps & {
| "outline" | "outline"
| "secondary" | "secondary"
| "ghost" | "ghost"
| "link"; | "link"
| "brand";
}; };
export const TooltipIconButton = forwardRef< export const TooltipIconButton = forwardRef<

View file

@ -1,13 +1,7 @@
"use client"; "use client";
import { Button } from "@/components/ui/button"; import { Button } from "@/components/ui/button";
import { import { Card, CardContent } from "@/components/ui/card";
Card, import { FilePlus2, Archive, Zap } from "lucide-react";
CardContent,
CardDescription,
CardHeader,
CardTitle,
} from "@/components/ui/card";
import { FilePlus2, Archive } from "lucide-react";
import { useRouter } from "next/navigation"; import { useRouter } from "next/navigation";
import { ThreadDisplayInfo } from "./types"; import { ThreadDisplayInfo } from "./types";
import { TerminalInput } from "./terminal-input"; import { TerminalInput } from "./terminal-input";
@ -28,6 +22,7 @@ import { QuickActions } from "./quick-actions";
import { useState } from "react"; import { useState } from "react";
import { GitHubLogoutButton } from "../github/github-oauth-button"; import { GitHubLogoutButton } from "../github/github-oauth-button";
import { MANAGER_GRAPH_ID } from "@open-swe/shared/constants"; import { MANAGER_GRAPH_ID } from "@open-swe/shared/constants";
import { TooltipIconButton } from "../ui/tooltip-icon-button";
interface DefaultViewProps { interface DefaultViewProps {
threads: ThreadDisplayInfo[]; threads: ThreadDisplayInfo[];
@ -48,6 +43,7 @@ export function DefaultView({ threads, threadsLoading }: DefaultViewProps) {
dragOver, dragOver,
handlePaste, handlePaste,
} = useFileUpload(); } = useFileUpload();
const [autoAccept, setAutoAccept] = useState(false);
if (!apiUrl) { if (!apiUrl) {
return <div>Missing API URL environment variable</div>; return <div>Missing API URL environment variable</div>;
@ -112,8 +108,10 @@ export function DefaultView({ threads, threadsLoading }: DefaultViewProps) {
onPaste={handlePaste} onPaste={handlePaste}
quickActionPrompt={quickActionPrompt} quickActionPrompt={quickActionPrompt}
setQuickActionPrompt={setQuickActionPrompt} setQuickActionPrompt={setQuickActionPrompt}
autoAcceptPlan={autoAccept}
setAutoAcceptPlan={setAutoAccept}
/> />
<div className="flex items-center gap-1"> <div className="flex items-center gap-2">
<TooltipProvider> <TooltipProvider>
<Tooltip> <Tooltip>
<TooltipTrigger> <TooltipTrigger>
@ -124,9 +122,24 @@ export function DefaultView({ threads, threadsLoading }: DefaultViewProps) {
<FilePlus2 className="size-4" /> <FilePlus2 className="size-4" />
</Label> </Label>
</TooltipTrigger> </TooltipTrigger>
<TooltipContent>Attach files</TooltipContent> <TooltipContent side="bottom">
Attach files
</TooltipContent>
</Tooltip> </Tooltip>
</TooltipProvider> </TooltipProvider>
<TooltipIconButton
variant={autoAccept ? "brand" : "ghost"}
tooltip="Automatically accept the plan"
className={cn(
autoAccept
? "text-secondary"
: "text-muted-foreground hover:text-foreground",
)}
onClick={() => setAutoAccept((prev) => !prev)}
side="bottom"
>
<Zap className="size-4" />
</TooltipIconButton>
</div> </div>
</div> </div>
</CardContent> </CardContent>

View file

@ -15,6 +15,7 @@ import { Base64ContentBlock, HumanMessage } from "@langchain/core/messages";
import { toast } from "sonner"; import { toast } from "sonner";
import { DEFAULT_CONFIG_KEY, useConfigStore } from "@/hooks/useConfigStore"; import { DEFAULT_CONFIG_KEY, useConfigStore } from "@/hooks/useConfigStore";
import { MANAGER_GRAPH_ID } from "@open-swe/shared/constants"; import { MANAGER_GRAPH_ID } from "@open-swe/shared/constants";
import { ManagerGraphUpdate } from "@open-swe/shared/open-swe/manager/types";
interface TerminalInputProps { interface TerminalInputProps {
placeholder?: string; placeholder?: string;
@ -26,6 +27,8 @@ interface TerminalInputProps {
onPaste?: (e: React.ClipboardEvent<HTMLTextAreaElement>) => void; onPaste?: (e: React.ClipboardEvent<HTMLTextAreaElement>) => void;
quickActionPrompt?: string; quickActionPrompt?: string;
setQuickActionPrompt?: Dispatch<SetStateAction<string>>; setQuickActionPrompt?: Dispatch<SetStateAction<string>>;
autoAcceptPlan: boolean;
setAutoAcceptPlan: Dispatch<SetStateAction<boolean>>;
} }
export function TerminalInput({ export function TerminalInput({
@ -38,6 +41,8 @@ export function TerminalInput({
onPaste, onPaste,
quickActionPrompt, quickActionPrompt,
setQuickActionPrompt, setQuickActionPrompt,
autoAcceptPlan,
setAutoAcceptPlan,
}: TerminalInputProps) { }: TerminalInputProps) {
const { push } = useRouter(); const { push } = useRouter();
const [message, setMessage] = useState(""); const [message, setMessage] = useState("");
@ -52,7 +57,6 @@ export function TerminalInput({
}); });
const handleSend = async () => { const handleSend = async () => {
const assistantId = MANAGER_GRAPH_ID;
if (!selectedRepository) { if (!selectedRepository) {
toast.error("Please select a repository first", { toast.error("Please select a repository first", {
richColors: true, richColors: true,
@ -75,27 +79,34 @@ export function TerminalInput({
try { try {
const newThreadId = uuidv4(); const newThreadId = uuidv4();
const run = await stream.client.runs.create(newThreadId, assistantId, { const runInput: ManagerGraphUpdate = {
input: { messages: [newHumanMessage],
messages: [newHumanMessage], targetRepository: selectedRepository,
targetRepository: selectedRepository, autoAcceptPlan,
}, };
config: { const run = await stream.client.runs.create(
recursion_limit: 400, newThreadId,
configurable: { MANAGER_GRAPH_ID,
...getConfig(DEFAULT_CONFIG_KEY), {
input: runInput,
config: {
recursion_limit: 400,
configurable: {
...getConfig(DEFAULT_CONFIG_KEY),
},
}, },
ifNotExists: "create",
streamResumable: true,
streamMode: ["values", "messages", "custom"],
}, },
ifNotExists: "create", );
streamResumable: true,
streamMode: ["values", "messages", "custom"],
});
// set session storage so the stream can be resumed after redirect. // set session storage so the stream can be resumed after redirect.
sessionStorage.setItem(`lg:stream:${newThreadId}`, run.run_id); sessionStorage.setItem(`lg:stream:${newThreadId}`, run.run_id);
push(`/chat/${newThreadId}`); push(`/chat/${newThreadId}`);
setMessage(""); setMessage("");
setContentBlocks([]); setContentBlocks([]);
setAutoAcceptPlan(false);
} catch (e) { } catch (e) {
console.error(e); console.error(e);
} finally { } finally {
@ -121,36 +132,25 @@ export function TerminalInput({
return ( return (
<div className="border-border bg-muted rounded-md border p-2 font-mono text-xs dark:bg-black"> <div className="border-border bg-muted rounded-md border p-2 font-mono text-xs dark:bg-black">
<div className="text-foreground flex items-start gap-1"> <div className="text-foreground flex items-center gap-1">
<span className="text-muted-foreground">open-swe</span> <div className="flex items-center gap-1 rounded-md border border-gray-200 p-1">
<span className="text-muted-foreground/70">@</span> <span className="text-muted-foreground">open-swe</span>
<span className="text-muted-foreground">github</span> <span className="text-muted-foreground/70">@</span>
<span className="text-muted-foreground/70">:</span> <span className="text-muted-foreground">github</span>
</div>
{/* Repository & Branch Selectors */} {/* Repository & Branch Selectors */}
<RepositoryBranchSelectors /> <RepositoryBranchSelectors />
{/* Prompt */} {/* Prompt */}
<span className="text-muted-foreground">$</span> <span className="text-muted-foreground">$</span>
</div>
{/* Multiline Input */}
<div className="mt-1 flex gap-2">
<Textarea
value={message}
onChange={(e) => setMessage(e.target.value)}
onKeyDown={handleKeyPress}
placeholder={placeholder}
disabled={disabled}
className="text-foreground placeholder:text-muted-foreground min-h-[40px] flex-1 resize-none border-none bg-transparent p-0 font-mono text-xs shadow-none focus-visible:ring-0 focus-visible:ring-offset-0"
rows={3}
onPaste={onPaste}
/>
<Button <Button
onClick={handleSend} onClick={handleSend}
disabled={disabled || !message.trim() || !selectedRepository} disabled={disabled || !message.trim() || !selectedRepository}
size="icon" size="icon"
variant="brand" variant="brand"
className="ml-auto size-8"
> >
{loading ? ( {loading ? (
<Loader2 className="size-4 animate-spin" /> <Loader2 className="size-4 animate-spin" />
@ -160,6 +160,20 @@ export function TerminalInput({
</Button> </Button>
</div> </div>
{/* Multiline Input */}
<div className="mt-2 flex gap-2">
<Textarea
value={message}
onChange={(e) => setMessage(e.target.value)}
onKeyDown={handleKeyPress}
placeholder={placeholder}
disabled={disabled}
className="text-foreground placeholder:text-muted-foreground min-h-[80px] flex-1 resize-none border-none bg-transparent p-0 font-mono text-xs shadow-none focus-visible:ring-0 focus-visible:ring-offset-0"
rows={6}
onPaste={onPaste}
/>
</div>
{/* Help text */} {/* Help text */}
<div className="text-muted-foreground mt-1 text-xs"> <div className="text-muted-foreground mt-1 text-xs">
Press Cmd+Enter to send Press Cmd+Enter to send

View file

@ -1,6 +1,7 @@
import { MessagesZodState } from "@langchain/langgraph"; import { MessagesZodState } from "@langchain/langgraph";
import { TargetRepository, TaskPlan, AgentSession } from "../types.js"; import { TargetRepository, TaskPlan, AgentSession } from "../types.js";
import { z } from "zod"; import { z } from "zod";
import { withLangGraph } from "@langchain/langgraph/zod";
export const ManagerGraphStateObj = MessagesZodState.extend({ export const ManagerGraphStateObj = MessagesZodState.extend({
/** /**
@ -34,6 +35,15 @@ export const ManagerGraphStateObj = MessagesZodState.extend({
* Can be user specified, or defaults to `open-swe/<manager-thread-id> * Can be user specified, or defaults to `open-swe/<manager-thread-id>
*/ */
branchName: z.string(), branchName: z.string(),
/**
* Whether or not to auto accept the generated plan.
*/
autoAcceptPlan: withLangGraph(z.custom<boolean>().optional(), {
reducer: {
schema: z.custom<boolean>().optional(),
fn: (_state, update) => update,
},
}),
}); });
export type ManagerGraphState = z.infer<typeof ManagerGraphStateObj>; export type ManagerGraphState = z.infer<typeof ManagerGraphStateObj>;

View file

@ -85,6 +85,12 @@ export const PlannerGraphStateObj = MessagesZodState.extend({
fn: (_state, update) => update, fn: (_state, update) => update,
}, },
}), }),
autoAcceptPlan: withLangGraph(z.custom<boolean>().optional(), {
reducer: {
schema: z.custom<boolean>().optional(),
fn: (_state, update) => update,
},
}),
}); });
export type PlannerGraphState = z.infer<typeof PlannerGraphStateObj>; export type PlannerGraphState = z.infer<typeof PlannerGraphStateObj>;