feat: Add restart from last checkpoint button (#724)

This commit is contained in:
Brace Sproul 2025-08-20 10:40:50 -07:00 • committed by GitHub
parent 81df6ffebf
commit df43631d59
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 342 additions and 18 deletions

View file

@ -26,7 +26,7 @@ import {
} from "../../../utils/github/issue-messages.js";
import { getBranchName } from "../../../utils/github/git.js";
import { getDefaultHeaders } from "../../../utils/default-headers.js";
import { getCustomConfigurableFields } from "../../../utils/config.js";
import { getCustomConfigurableFields } from "@open-swe/shared/open-swe/utils/config";
import { StreamMode } from "@langchain/langgraph-sdk";
import { isLocalMode } from "@open-swe/shared/open-swe/local-mode";
import { regenerateInstallationToken } from "../../../utils/github/regenerate-token.js";

View file

@ -18,7 +18,7 @@ import { createLogger, LogLevel } from "../../../utils/logger.js";
import { getBranchName } from "../../../utils/github/git.js";
import { PlannerGraphUpdate } from "@open-swe/shared/open-swe/planner/types";
import { getDefaultHeaders } from "../../../utils/default-headers.js";
import { getCustomConfigurableFields } from "../../../utils/config.js";
import { getCustomConfigurableFields } from "@open-swe/shared/open-swe/utils/config";
import { getRecentUserRequest } from "../../../utils/user-request.js";
import { StreamMode } from "@langchain/langgraph-sdk";
import { regenerateInstallationToken } from "../../../utils/github/regenerate-token.js";

View file

@ -42,7 +42,7 @@ import {
CustomNodeEvent,
} from "@open-swe/shared/open-swe/custom-node-events";
import { getDefaultHeaders } from "../../../utils/default-headers.js";
import { getCustomConfigurableFields } from "../../../utils/config.js";
import { getCustomConfigurableFields } from "@open-swe/shared/open-swe/utils/config";
import { isLocalMode } from "@open-swe/shared/open-swe/local-mode";
import {
postGitHubIssueComment,

View file

@ -0,0 +1,169 @@
import { v4 as uuidv4 } from "uuid";
import {
GITHUB_TOKEN_COOKIE,
GITHUB_INSTALLATION_ID_COOKIE,
GITHUB_INSTALLATION_TOKEN_COOKIE,
GITHUB_INSTALLATION_NAME,
GITHUB_INSTALLATION_ID,
PROGRAMMER_GRAPH_ID,
PLANNER_GRAPH_ID,
OPEN_SWE_STREAM_MODE,
MANAGER_GRAPH_ID,
} from "@open-swe/shared/constants";
import {
getGitHubInstallationTokenOrThrow,
getInstallationNameFromReq,
getGitHubAccessTokenOrThrow,
} from "../[..._path]/utils";
import { NextRequest, NextResponse } from "next/server";
import { RestartRunRequest } from "./types";
import { Client, StreamMode, ThreadState } from "@langchain/langgraph-sdk";
import { ManagerGraphState } from "@open-swe/shared/open-swe/manager/types";
import { PlannerGraphState } from "@open-swe/shared/open-swe/planner/types";
import { AgentSession, GraphState } from "@open-swe/shared/open-swe/types";
import { END } from "@langchain/langgraph/web";
import { getCustomConfigurableFields } from "@open-swe/shared/open-swe/utils/config";
async function getRequestHeaders(
req: NextRequest,
): Promise<Record<string, string>> {
const encryptionKey = process.env.SECRETS_ENCRYPTION_KEY;
if (!encryptionKey) {
throw new Error("SECRETS_ENCRYPTION_KEY environment variable is required");
}
const installationIdCookie = req.cookies.get(
GITHUB_INSTALLATION_ID_COOKIE,
)?.value;
if (!installationIdCookie) {
throw new Error(
"No GitHub installation ID found. GitHub App must be installed first.",
);
}
const [installationToken, installationName] = await Promise.all([
getGitHubInstallationTokenOrThrow(installationIdCookie, encryptionKey),
getInstallationNameFromReq(req, installationIdCookie),
]);
return {
[GITHUB_TOKEN_COOKIE]: getGitHubAccessTokenOrThrow(req, encryptionKey),
[GITHUB_INSTALLATION_TOKEN_COOKIE]: installationToken,
[GITHUB_INSTALLATION_NAME]: installationName,
[GITHUB_INSTALLATION_ID]: installationIdCookie,
};
}
async function createNewSession(
client: Client,
inputs: {
graphId: string;
threadState: ThreadState<
ManagerGraphState | PlannerGraphState | GraphState
>;
},
): Promise<AgentSession> {
const newThreadId = uuidv4();
const hasNext = inputs.threadState.next.length > 0;
const run = await client.runs.create(newThreadId, inputs.graphId, {
command: {
update: inputs.threadState.values,
...(hasNext ? { goto: inputs.threadState.next[0] } : { goto: END }),
},
ifNotExists: "create",
streamMode: OPEN_SWE_STREAM_MODE as StreamMode[],
streamResumable: true,
config: {
recursion_limit: 400,
configurable: getCustomConfigurableFields(
inputs.threadState.metadata as Record<string, any>,
),
},
});
return {
threadId: newThreadId,
runId: run.run_id,
};
}
/**
* Restart a run. This function isn't actually restarting a run,
* but rather it's creating fresh new threads & runs for all existing
* threads. Whichever thread was the one to fail will be restarted and
* resumed where it was failed.
*/
export async function POST(request: NextRequest): Promise<NextResponse> {
try {
const body: RestartRunRequest = await request.json();
const { managerThreadId, plannerThreadId, programmerThreadId } = body;
const langGraphClient = new Client({
apiUrl: process.env.LANGGRAPH_API_URL ?? "http://localhost:2024",
defaultHeaders: await getRequestHeaders(request),
});
const [managerThreadState, plannerThreadState, programmerThreadState] =
await Promise.all([
langGraphClient.threads.getState<ManagerGraphState>(managerThreadId),
langGraphClient.threads.getState<PlannerGraphState>(plannerThreadId),
programmerThreadId
? langGraphClient.threads.getState<GraphState>(programmerThreadId)
: null,
]);
if (!managerThreadState || !plannerThreadState) {
return NextResponse.json(
{
error:
"Failed to restart run. Must have existing planner and manager threads.",
},
{ status: 500 },
);
}
const newProgrammerSession = programmerThreadState
? await createNewSession(langGraphClient, {
graphId: PROGRAMMER_GRAPH_ID,
threadState: programmerThreadState,
})
: undefined;
const newPlannerState: PlannerGraphState = {
...plannerThreadState.values,
...(newProgrammerSession
? {
programmerSession: newProgrammerSession,
}
: {}),
};
const newPlannerSession = await createNewSession(langGraphClient, {
graphId: PLANNER_GRAPH_ID,
threadState: {
...plannerThreadState,
values: newPlannerState,
},
});
const newManagerState: ManagerGraphState = {
...managerThreadState.values,
plannerSession: newPlannerSession,
};
const newManagerSession = await createNewSession(langGraphClient, {
graphId: MANAGER_GRAPH_ID,
threadState: {
...managerThreadState,
values: newManagerState,
},
});
return NextResponse.json({
managerSession: newManagerSession,
plannerSession: newPlannerSession,
programmerSession: newProgrammerSession,
});
} catch (error) {
console.error("Failed to restart run", error);
return NextResponse.json(
{ error: "Failed to restart run" },
{ status: 500 },
);
}
}

View file

@ -0,0 +1,14 @@
import { AgentSession } from "@open-swe/shared/open-swe/types";
export interface RestartRunRequest {
managerThreadId: string;
plannerThreadId: string;
// Programmer thread ID can be undefined if the error occurred in the planner graph
programmerThreadId?: string;
}
export interface RestartRunResponse {
managerSession: AgentSession;
plannerSession: AgentSession;
programmerSession: AgentSession;
}

View file

@ -14,6 +14,7 @@ import { ErrorState } from "./types";
import { CollapsibleAlert } from "./collapsible-alert";
import { Loader2 } from "lucide-react";
import { parsePartialJson } from "@langchain/core/output_parsers";
import { RestartRun } from "./restart-run";
function MessageCopyButton({ content }: { content: string }) {
const [copied, setCopied] = useState(false);
@ -70,6 +71,11 @@ interface ManagerChatProps {
isLoading: boolean;
cancelRun: () => void;
errorState?: ErrorState | null;
// Restart-run controls
canRestartRun?: boolean;
managerThreadId?: string;
plannerThreadId?: string;
programmerThreadId?: string;
githubUser?: {
login: string;
avatar_url: string;
@ -162,6 +168,10 @@ export function ManagerChat({
isLoading,
cancelRun,
errorState,
canRestartRun,
managerThreadId,
plannerThreadId,
programmerThreadId,
githubUser,
disableSubmit,
}: ManagerChatProps) {
@ -235,6 +245,13 @@ export function ManagerChat({
icon={<AlertCircle className="size-4" />}
/>
) : null}
{canRestartRun && managerThreadId && plannerThreadId ? (
<RestartRun
managerThreadId={managerThreadId}
plannerThreadId={plannerThreadId}
programmerThreadId={programmerThreadId}
/>
) : null}
</>
}
footer={

View file

@ -0,0 +1,135 @@
"use client";
import { useState } from "react";
import { Button } from "@/components/ui/button";
import { cn } from "@/lib/utils";
import {
Loader2,
RotateCcw,
CheckCircle2,
AlertTriangle,
ExternalLink,
} from "lucide-react";
import type {
RestartRunRequest,
RestartRunResponse,
} from "@/app/api/restart-run/types";
import Link from "next/link";
interface RestartRunProps {
managerThreadId: string;
plannerThreadId: string;
programmerThreadId?: string;
className?: string;
}
export function RestartRun({
managerThreadId,
plannerThreadId,
programmerThreadId,
className,
}: RestartRunProps) {
const [isLoading, setIsLoading] = useState(false);
const [error, setError] = useState<string | null>(null);
const [response, setResponse] = useState<RestartRunResponse | null>(null);
const handleRestart = async (): Promise<void> => {
setIsLoading(true);
setError(null);
setResponse(null);
const body: RestartRunRequest = {
managerThreadId,
plannerThreadId,
...(programmerThreadId ? { programmerThreadId } : {}),
};
try {
const res = await fetch("/api/restart-run", {
method: "POST",
headers: { "Content-Type": "application/json" },
body: JSON.stringify(body),
});
if (!res.ok) {
const data = (await res.json().catch(() => ({}))) as {
error?: string;
};
throw new Error(data.error || "Failed to restart run");
}
const data = (await res.json()) as RestartRunResponse;
setResponse(data);
} catch (e) {
const message =
e instanceof Error ? e.message : "Unexpected error restarting run";
setError(message);
} finally {
setIsLoading(false);
}
};
return (
<div
className={cn(
"border-border/60 from-background to-background/95 rounded-lg border bg-gradient-to-r p-4",
className,
)}
>
{response ? (
<div className="space-y-3">
<div className="flex items-center gap-2">
<div className="flex size-6 items-center justify-center rounded-full bg-green-100 dark:bg-green-900/30">
<CheckCircle2 className="size-3 text-green-600 dark:text-green-400" />
</div>
<span className="text-foreground text-sm font-medium">
Run created successfully
</span>
</div>
<Link
href={`/chat/${response.managerSession.threadId}`}
className="group bg-primary/10 text-primary hover:bg-primary/20 inline-flex items-center gap-2 rounded-md px-3 py-1.5 text-sm transition-colors"
>
Open session
<ExternalLink className="size-3 transition-transform group-hover:translate-x-0.5" />
</Link>
</div>
) : (
<div className="space-y-3">
<div className="flex items-center gap-2">
<div className="flex size-6 items-center justify-center rounded-full bg-red-100 dark:bg-red-900/30">
<AlertTriangle className="size-3 text-red-600 dark:text-red-400" />
</div>
<span className="text-sm font-medium text-red-700 dark:text-red-300">
A fatal error occurred
</span>
</div>
<div className="flex items-center gap-3">
<Button
onClick={handleRestart}
disabled={isLoading}
size="sm"
variant="outline"
className="h-7 px-3 text-xs"
>
{isLoading ? (
<>
<Loader2 className="mr-1.5 size-3 animate-spin" />
Restarting...
</>
) : (
<>
<RotateCcw className="mr-1.5 size-3" />
Restart from last checkpoint
</>
)}
</Button>
{error && <span className="text-destructive text-xs">{error}</span>}
</div>
</div>
)}
</div>
);
}

View file

@ -408,6 +408,10 @@ export function ThreadView({
isLoading={stream.isLoading}
cancelRun={cancelRun}
errorState={errorState}
canRestartRun={Boolean(plannerStream.error || programmerStream.error)}
managerThreadId={displayThread.id}
plannerThreadId={plannerSession?.threadId}
programmerThreadId={programmerSession?.threadId}
githubUser={user || undefined}
disableSubmit={shouldDisableManagerInput}
/>

View file

@ -8,11 +8,6 @@ import {
} from "@langchain/langgraph/web";
import { MODEL_OPTIONS, MODEL_OPTIONS_NO_THINKING } from "./models.js";
import { ConfigurableFieldUIMetadata } from "../configurable-metadata.js";
import {
uiMessageReducer,
type UIMessage,
type RemoveUIMessage,
} from "@langchain/langgraph-sdk/react-ui";
import {
GITHUB_INSTALLATION_NAME,
GITHUB_INSTALLATION_TOKEN_COOKIE,
@ -289,16 +284,6 @@ export const GraphAnnotation = MessagesZodState.extend({
fn: tokenDataReducer,
},
}),
// ---NOT USED---
ui: z
.custom<UIMessage[]>()
.default(() => [])
.langgraph.reducer<(UIMessage | RemoveUIMessage)[]>((state, update) =>
uiMessageReducer(state, update),
),
// TODO: Not used, but can be used in the future for Gen UI artifacts
context: z.record(z.string(), z.unknown()),
});
export type GraphState = z.infer<typeof GraphAnnotation>;