mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-01 09:43:14 +00:00
feat: Add restart from last checkpoint button (#724)
This commit is contained in:
parent
81df6ffebf
commit
df43631d59
10 changed files with 342 additions and 18 deletions
|
|
@ -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";
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
169
apps/web/src/app/api/restart-run/route.ts
Normal file
169
apps/web/src/app/api/restart-run/route.ts
Normal 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 },
|
||||
);
|
||||
}
|
||||
}
|
||||
14
apps/web/src/app/api/restart-run/types.ts
Normal file
14
apps/web/src/app/api/restart-run/types.ts
Normal 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;
|
||||
}
|
||||
|
|
@ -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={
|
||||
|
|
|
|||
135
apps/web/src/components/v2/restart-run.tsx
Normal file
135
apps/web/src/components/v2/restart-run.tsx
Normal 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>
|
||||
);
|
||||
}
|
||||
|
|
@ -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}
|
||||
/>
|
||||
|
|
|
|||
|
|
@ -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>;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue