mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-06 15:52:11 +00:00
fix: Pass proper config when retrying run (#808)
* fix: Pass proper config when retrying run * fix
This commit is contained in:
parent
504be7c563
commit
d81080e71e
1 changed files with 30 additions and 12 deletions
|
|
@ -20,7 +20,11 @@ import { RestartRunRequest } from "./types";
|
||||||
import { Client, StreamMode, ThreadState } from "@langchain/langgraph-sdk";
|
import { Client, StreamMode, ThreadState } from "@langchain/langgraph-sdk";
|
||||||
import { ManagerGraphState } from "@openswe/shared/open-swe/manager/types";
|
import { ManagerGraphState } from "@openswe/shared/open-swe/manager/types";
|
||||||
import { PlannerGraphState } from "@openswe/shared/open-swe/planner/types";
|
import { PlannerGraphState } from "@openswe/shared/open-swe/planner/types";
|
||||||
import { AgentSession, GraphState } from "@openswe/shared/open-swe/types";
|
import {
|
||||||
|
AgentSession,
|
||||||
|
GraphConfig,
|
||||||
|
GraphState,
|
||||||
|
} from "@openswe/shared/open-swe/types";
|
||||||
import { END } from "@langchain/langgraph/web";
|
import { END } from "@langchain/langgraph/web";
|
||||||
import { getCustomConfigurableFields } from "@openswe/shared/open-swe/utils/config";
|
import { getCustomConfigurableFields } from "@openswe/shared/open-swe/utils/config";
|
||||||
|
|
||||||
|
|
@ -60,10 +64,12 @@ async function createNewSession(
|
||||||
threadState: ThreadState<
|
threadState: ThreadState<
|
||||||
ManagerGraphState | PlannerGraphState | GraphState
|
ManagerGraphState | PlannerGraphState | GraphState
|
||||||
>;
|
>;
|
||||||
|
threadConfig: GraphConfig;
|
||||||
},
|
},
|
||||||
): Promise<AgentSession> {
|
): Promise<AgentSession> {
|
||||||
const newThreadId = uuidv4();
|
const newThreadId = uuidv4();
|
||||||
const hasNext = inputs.threadState.next.length > 0;
|
const hasNext = inputs.threadState.next.length > 0;
|
||||||
|
|
||||||
const run = await client.runs.create(newThreadId, inputs.graphId, {
|
const run = await client.runs.create(newThreadId, inputs.graphId, {
|
||||||
command: {
|
command: {
|
||||||
update: inputs.threadState.values,
|
update: inputs.threadState.values,
|
||||||
|
|
@ -74,9 +80,7 @@ async function createNewSession(
|
||||||
streamResumable: true,
|
streamResumable: true,
|
||||||
config: {
|
config: {
|
||||||
recursion_limit: 400,
|
recursion_limit: 400,
|
||||||
configurable: getCustomConfigurableFields(
|
configurable: getCustomConfigurableFields(inputs.threadConfig),
|
||||||
inputs.threadState.metadata as Record<string, any>,
|
|
||||||
),
|
|
||||||
},
|
},
|
||||||
});
|
});
|
||||||
return {
|
return {
|
||||||
|
|
@ -101,14 +105,25 @@ export async function POST(request: NextRequest): Promise<NextResponse> {
|
||||||
defaultHeaders: await getRequestHeaders(request),
|
defaultHeaders: await getRequestHeaders(request),
|
||||||
});
|
});
|
||||||
|
|
||||||
const [managerThreadState, plannerThreadState, programmerThreadState] =
|
const [
|
||||||
await Promise.all([
|
managerThread,
|
||||||
langGraphClient.threads.getState<ManagerGraphState>(managerThreadId),
|
managerThreadState,
|
||||||
langGraphClient.threads.getState<PlannerGraphState>(plannerThreadId),
|
plannerThread,
|
||||||
programmerThreadId
|
plannerThreadState,
|
||||||
? langGraphClient.threads.getState<GraphState>(programmerThreadId)
|
programmerThread,
|
||||||
: null,
|
programmerThreadState,
|
||||||
]);
|
] = await Promise.all([
|
||||||
|
langGraphClient.threads.get<ManagerGraphState>(managerThreadId),
|
||||||
|
langGraphClient.threads.getState<ManagerGraphState>(managerThreadId),
|
||||||
|
langGraphClient.threads.get<PlannerGraphState>(plannerThreadId),
|
||||||
|
langGraphClient.threads.getState<PlannerGraphState>(plannerThreadId),
|
||||||
|
programmerThreadId
|
||||||
|
? langGraphClient.threads.get<GraphState>(programmerThreadId)
|
||||||
|
: null,
|
||||||
|
programmerThreadId
|
||||||
|
? langGraphClient.threads.getState<GraphState>(programmerThreadId)
|
||||||
|
: null,
|
||||||
|
]);
|
||||||
if (!managerThreadState || !plannerThreadState) {
|
if (!managerThreadState || !plannerThreadState) {
|
||||||
return NextResponse.json(
|
return NextResponse.json(
|
||||||
{
|
{
|
||||||
|
|
@ -123,6 +138,7 @@ export async function POST(request: NextRequest): Promise<NextResponse> {
|
||||||
? await createNewSession(langGraphClient, {
|
? await createNewSession(langGraphClient, {
|
||||||
graphId: PROGRAMMER_GRAPH_ID,
|
graphId: PROGRAMMER_GRAPH_ID,
|
||||||
threadState: programmerThreadState,
|
threadState: programmerThreadState,
|
||||||
|
threadConfig: (programmerThread as Record<string, any>)?.config,
|
||||||
})
|
})
|
||||||
: undefined;
|
: undefined;
|
||||||
|
|
||||||
|
|
@ -140,6 +156,7 @@ export async function POST(request: NextRequest): Promise<NextResponse> {
|
||||||
...plannerThreadState,
|
...plannerThreadState,
|
||||||
values: newPlannerState,
|
values: newPlannerState,
|
||||||
},
|
},
|
||||||
|
threadConfig: (plannerThread as Record<string, any>)?.config,
|
||||||
});
|
});
|
||||||
|
|
||||||
const newManagerState: ManagerGraphState = {
|
const newManagerState: ManagerGraphState = {
|
||||||
|
|
@ -152,6 +169,7 @@ export async function POST(request: NextRequest): Promise<NextResponse> {
|
||||||
...managerThreadState,
|
...managerThreadState,
|
||||||
values: newManagerState,
|
values: newManagerState,
|
||||||
},
|
},
|
||||||
|
threadConfig: (managerThread as Record<string, any>)?.config,
|
||||||
});
|
});
|
||||||
|
|
||||||
return NextResponse.json({
|
return NextResponse.json({
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue