fix: Properly join streams on client (#229)

* fix: Properly join streams on client

* format and lint
This commit is contained in:
Brace Sproul 2025-06-18 13:12:16 -07:00 • committed by GitHub
parent dbc9e00666
commit 6ef09d7396
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
22 changed files with 148 additions and 176 deletions

View file

@ -1,6 +1,6 @@
import { END, START, StateGraph } from "@langchain/langgraph";
import { GraphConfiguration } from "@open-swe/shared/open-swe/types";
import { ManagerGraphStateObj } from "./types.js";
import { ManagerGraphStateObj } from "@open-swe/shared/open-swe/manager/types";
import {
initializeGithubIssue,
classifyMessage,

View file

@ -1,5 +1,8 @@
import { GraphConfig, TaskPlan } from "@open-swe/shared/open-swe/types";
import { ManagerGraphState, ManagerGraphUpdate } from "../types.js";
import {
ManagerGraphState,
ManagerGraphUpdate,
} from "@open-swe/shared/open-swe/manager/types";
import { createLangGraphClient } from "../../../utils/langgraph-client.js";
import {
GITHUB_INSTALLATION_TOKEN_COOKIE,
@ -168,11 +171,11 @@ export async function classifyMessage(
});
const [programmerThread, plannerThread] = await Promise.all([
state.programmerThreadId
? langGraphClient.threads.get(state.programmerThreadId)
state.programmerSession?.threadId
? langGraphClient.threads.get(state.programmerSession.threadId)
: undefined,
state.plannerThreadId
? langGraphClient.threads.get(state.plannerThreadId)
state.plannerSession?.threadId
? langGraphClient.threads.get(state.plannerSession.threadId)
: undefined,
]);
const programmerStatus = programmerThread?.status ?? "not_started";

View file

@ -1,6 +1,9 @@
import { v4 as uuidv4 } from "uuid";
import { GraphConfig } from "@open-swe/shared/open-swe/types";
import { ManagerGraphState, ManagerGraphUpdate } from "../types.js";
import {
ManagerGraphState,
ManagerGraphUpdate,
} from "@open-swe/shared/open-swe/manager/types";
import { createIssueTitleAndBodyFromMessages } from "../utils/generate-issue-fields.js";
import {
GITHUB_INSTALLATION_TOKEN_COOKIE,
@ -92,6 +95,7 @@ ${ISSUE_CONTENT_CLOSE_TAG}`,
recursion_limit: 400,
},
ifNotExists: "create",
streamResumable: true,
});
return {

View file

@ -1,6 +1,9 @@
import { v4 as uuidv4 } from "uuid";
import { GraphConfig } from "@open-swe/shared/open-swe/types";
import { ManagerGraphState, ManagerGraphUpdate } from "../types.js";
import {
ManagerGraphState,
ManagerGraphUpdate,
} from "@open-swe/shared/open-swe/manager/types";
import { getGitHubTokensFromConfig } from "../../../utils/github-tokens.js";
import { HumanMessage, isHumanMessage } from "@langchain/core/messages";
import { getIssue } from "../../../utils/github/api.js";

View file

@ -1,6 +1,9 @@
import { v4 as uuidv4 } from "uuid";
import { GraphConfig } from "@open-swe/shared/open-swe/types";
import { ManagerGraphState } from "../types.js";
import {
ManagerGraphState,
ManagerGraphUpdate,
} from "@open-swe/shared/open-swe/manager/types";
import { createLangGraphClient } from "../../../utils/langgraph-client.js";
import {
GITHUB_INSTALLATION_TOKEN_COOKIE,
@ -18,7 +21,7 @@ const logger = createLogger(LogLevel.INFO, "StartPlanner");
export async function startPlanner(
state: ManagerGraphState,
config: GraphConfig,
) {
): Promise<ManagerGraphUpdate> {
const langGraphClient = createLangGraphClient({
defaultHeaders: {
[GITHUB_TOKEN_COOKIE]: config.configurable?.[GITHUB_TOKEN_COOKIE] ?? "",
@ -27,9 +30,9 @@ export async function startPlanner(
},
});
const plannerThreadId = state.plannerThreadId ?? uuidv4();
const plannerThreadId = state.plannerSession?.threadId ?? uuidv4();
try {
await langGraphClient.runs.create(plannerThreadId, "planner", {
const run = await langGraphClient.runs.create(plannerThreadId, "planner", {
input: {
// github issue ID & target repo so the planning agent can fetch the user's request, and clone the repo.
githubIssueId: state.githubIssueId,
@ -43,10 +46,14 @@ export async function startPlanner(
},
ifNotExists: "create",
multitaskStrategy: "enqueue",
streamResumable: true,
});
return {
plannerThreadId,
plannerSession: {
threadId: plannerThreadId,
runId: run.run_id,
},
};
} catch (error) {
logger.error("Failed to start planner", {

View file

@ -1,40 +0,0 @@
import { MessagesZodState } from "@langchain/langgraph";
import { TargetRepository, TaskPlan } from "@open-swe/shared/open-swe/types";
import { z } from "zod";
export const ManagerGraphStateObj = MessagesZodState.extend({
/**
* The GitHub issue number that the user's request is associated with.
* If not provided when the graph is invoked, it will create an issue.
*/
githubIssueId: z.number(),
/**
* The GitHub pull request number of the PR which resolves the user's request.
* If not provided when the graph is invoked, it will create a PR.
*/
githubPullRequestId: z.number().optional(),
/**
* The target repository the request should be executed in.
*/
targetRepository: z.custom<TargetRepository>(),
/**
* The tasks generated for this request.
*/
taskPlan: z.custom<TaskPlan>(),
/**
* The programmer thread ID
*/
programmerThreadId: z.string().optional(),
/**
* The planner thread ID
*/
plannerThreadId: z.string().optional(),
/**
* The branch name to checkout and make changes on.
* Can be user specified, or defaults to `open-swe/<manager-thread-id>
*/
branchName: z.string(),
});
export type ManagerGraphState = z.infer<typeof ManagerGraphStateObj>;
export type ManagerGraphUpdate = Partial<ManagerGraphState>;

View file

@ -1,5 +1,8 @@
import { END, START, StateGraph } from "@langchain/langgraph";
import { PlannerGraphState, PlannerGraphStateObj } from "./types.js";
import {
PlannerGraphState,
PlannerGraphStateObj,
} from "@open-swe/shared/open-swe/planner/types";
import {
GraphConfig,
GraphConfiguration,

View file

@ -1,6 +1,9 @@
import { loadModel, Task } from "../../../../utils/load-model.js";
import { createShellTool } from "../../../../tools/index.js";
import { PlannerGraphState, PlannerGraphUpdate } from "../../types.js";
import {
PlannerGraphState,
PlannerGraphUpdate,
} from "@open-swe/shared/open-swe/planner/types";
import { GraphConfig } from "@open-swe/shared/open-swe/types";
import { createLogger, LogLevel } from "../../../../utils/logger.js";
import { getMessageContentString } from "@open-swe/shared/messages";

View file

@ -2,7 +2,10 @@ import { isAIMessage, ToolMessage } from "@langchain/core/messages";
import { createSessionPlanToolFields } from "../../../tools/index.js";
import { GraphConfig } from "@open-swe/shared/open-swe/types";
import { loadModel, Task } from "../../../utils/load-model.js";
import { PlannerGraphState, PlannerGraphUpdate } from "../types.js";
import {
PlannerGraphState,
PlannerGraphUpdate,
} from "@open-swe/shared/open-swe/planner/types";
import { getUserRequest } from "../../../utils/user-request.js";
import {
formatFollowupMessagePrompt,

View file

@ -1,4 +1,7 @@
import { PlannerGraphState, PlannerGraphUpdate } from "../types.js";
import {
PlannerGraphState,
PlannerGraphUpdate,
} from "@open-swe/shared/open-swe/planner/types";
import { Command, END } from "@langchain/langgraph";
import { getGitHubTokensFromConfig } from "../../../utils/github-tokens.js";
import { getIssue, getIssueComments } from "../../../utils/github/api.js";

View file

@ -15,7 +15,10 @@ import {
PLAN_INTERRUPT_ACTION_TITLE,
PLAN_INTERRUPT_DELIMITER,
} from "@open-swe/shared/constants";
import { PlannerGraphState, PlannerGraphUpdate } from "../types.js";
import {
PlannerGraphState,
PlannerGraphUpdate,
} from "@open-swe/shared/open-swe/planner/types";
import { createLangGraphClient } from "../../../utils/langgraph-client.js";
import { addTaskPlanToIssue } from "../../../utils/github/issue-task.js";
@ -109,13 +112,18 @@ export async function interruptProposedPlan(
// Restart the sandbox.
runInput.sandboxSessionId = (await startSandbox(state.sandboxSessionId)).id;
await langGraphClient.runs.create(programmerThreadId, "programmer", {
input: runInput,
config: {
recursion_limit: 400,
const run = await langGraphClient.runs.create(
programmerThreadId,
"programmer",
{
input: runInput,
config: {
recursion_limit: 400,
},
ifNotExists: "create",
streamResumable: true,
},
ifNotExists: "create",
});
);
await addTaskPlanToIssue(
{
@ -127,7 +135,10 @@ export async function interruptProposedPlan(
);
return {
programmerThreadId,
programmerSession: {
threadId: programmerThreadId,
runId: run.run_id,
},
sandboxSessionId: runInput.sandboxSessionId,
taskPlan: runInput.taskPlan,
};

View file

@ -6,7 +6,10 @@ import { z } from "zod";
import { tool } from "@langchain/core/tools";
import { ConfigurableModel } from "langchain/chat_models/universal";
import { traceable } from "langsmith/traceable";
import { PlannerGraphState, PlannerGraphUpdate } from "../types.js";
import {
PlannerGraphState,
PlannerGraphUpdate,
} from "@open-swe/shared/open-swe/planner/types";
import { getUserRequest } from "../../../utils/user-request.js";
import { loadModel, Task } from "../../../utils/load-model.js";

View file

@ -1,6 +1,9 @@
import { z } from "zod";
import { GraphConfig } from "@open-swe/shared/open-swe/types";
import { PlannerGraphState, PlannerGraphUpdate } from "../types.js";
import {
PlannerGraphState,
PlannerGraphUpdate,
} from "@open-swe/shared/open-swe/planner/types";
import { loadModel, Task } from "../../../utils/load-model.js";
import { getMessageString } from "../../../utils/message/content.js";
import { getUserRequest } from "../../../utils/user-request.js";

View file

@ -1,7 +1,10 @@
import { isAIMessage, ToolMessage } from "@langchain/core/messages";
import { createShellTool } from "../../../tools/index.js";
import { GraphConfig } from "@open-swe/shared/open-swe/types";
import { PlannerGraphState, PlannerGraphUpdate } from "../types.js";
import {
PlannerGraphState,
PlannerGraphUpdate,
} from "@open-swe/shared/open-swe/planner/types";
import { createLogger, LogLevel } from "../../../utils/logger.js";
import { zodSchemaToString } from "../../../utils/zod-to-string.js";
import { formatBadArgsError } from "../../../utils/zod-to-string.js";

View file

@ -1,73 +0,0 @@
import "@langchain/langgraph/zod";
import { z } from "zod";
import { MessagesZodState } from "@langchain/langgraph";
import { TargetRepository, TaskPlan } from "@open-swe/shared/open-swe/types";
import { withLangGraph } from "@langchain/langgraph/zod";
export const PlannerGraphStateObj = MessagesZodState.extend({
sandboxSessionId: withLangGraph(z.custom<string>(), {
reducer: {
schema: z.custom<string>(),
fn: (_state, update) => update,
},
}),
targetRepository: withLangGraph(z.custom<TargetRepository>(), {
reducer: {
schema: z.custom<TargetRepository>(),
fn: (_state, update) => update,
},
}),
githubIssueId: withLangGraph(z.custom<number>(), {
reducer: {
schema: z.custom<number>(),
fn: (_state, update) => update,
},
}),
codebaseTree: withLangGraph(z.custom<string>(), {
reducer: {
schema: z.custom<string>(),
fn: (_state, update) => update,
},
}),
taskPlan: withLangGraph(z.custom<TaskPlan>(), {
reducer: {
schema: z.custom<TaskPlan>(),
fn: (_state, update) => update,
},
}),
proposedPlan: withLangGraph(z.custom<string[]>(), {
reducer: {
schema: z.custom<string[]>(),
fn: (_state, update) => update,
},
default: (): string[] => [],
}),
planContextSummary: withLangGraph(z.custom<string>(), {
reducer: {
schema: z.custom<string>(),
fn: (_state, update) => update,
},
default: () => "",
}),
branchName: withLangGraph(z.custom<string>(), {
reducer: {
schema: z.custom<string>(),
fn: (_state, update) => update,
},
}),
planChangeRequest: withLangGraph(z.custom<string>(), {
reducer: {
schema: z.custom<string>(),
fn: (_state, update) => update,
},
}),
programmerThreadId: withLangGraph(z.custom<string>(), {
reducer: {
schema: z.custom<string>(),
fn: (_state, update) => update,
},
}),
});
export type PlannerGraphState = z.infer<typeof PlannerGraphStateObj>;
export type PlannerGraphUpdate = Partial<PlannerGraphState>;

View file

@ -47,7 +47,7 @@
"framer-motion": "^12.4.9",
"jsonwebtoken": "^9.0.2",
"katex": "^0.16.21",
"langgraph-nextjs-api-passthrough": "^0.1.2",
"langgraph-nextjs-api-passthrough": "^0.1.3",
"lodash": "^4.17.21",
"lucide-react": "^0.476.0",
"next-themes": "^0.4.4",

View file

@ -1,20 +1,25 @@
import { isHumanMessageSDK } from "@/lib/langchain-messages";
import { UseStream, useStream } from "@langchain/langgraph-sdk/react";
import { AssistantMessage } from "../thread/messages/ai";
import { useEffect } from "react";
import { useEffect, useRef } from "react";
import { ManagerGraphState } from "@open-swe/shared/open-swe/manager/types";
interface ActionsRendererProps {
graphId: string;
threadId: string;
setProgrammerThreadId?: (threadId: string) => void;
programmerThreadId?: string;
runId?: string;
setProgrammerSession?: (
session: ManagerGraphState["programmerSession"],
) => void;
programmerSession?: ManagerGraphState["programmerSession"];
}
export function ActionsRenderer<State extends Record<string, unknown>>({
graphId,
threadId,
setProgrammerThreadId,
programmerThreadId,
runId,
setProgrammerSession,
programmerSession,
}: ActionsRendererProps) {
const stream = useStream<State>({
apiUrl: process.env.NEXT_PUBLIC_API_URL,
@ -23,6 +28,14 @@ export function ActionsRenderer<State extends Record<string, unknown>>({
threadId,
});
const streamJoined = useRef(false);
useEffect(() => {
if (!streamJoined.current && runId) {
streamJoined.current = true;
stream.joinStream(runId).catch(console.error);
}
}, [runId]);
const nonHumanMessages = stream.messages?.filter(
(m) => !isHumanMessageSDK(m),
);
@ -30,11 +43,22 @@ export function ActionsRenderer<State extends Record<string, unknown>>({
// TODO: Need a better way to handle this. Not great like this...
useEffect(() => {
if (
stream.values?.programmerThreadId &&
typeof stream.values.programmerThreadId === "string" &&
!programmerThreadId
stream.values?.programmerSession &&
typeof stream.values.programmerSession === "object" &&
stream.values.programmerSession &&
(
stream.values
.programmerSession as ManagerGraphState["programmerSession"]
)?.runId &&
(
stream.values
.programmerSession as ManagerGraphState["programmerSession"]
)?.threadId &&
!programmerSession
) {
setProgrammerThreadId?.(stream.values.programmerThreadId as string);
const programmerSession = stream.values
.programmerSession as ManagerGraphState["programmerSession"];
setProgrammerSession?.(programmerSession);
}
}, [stream.values]);

View file

@ -33,8 +33,11 @@ export function ThreadView({
onBackToHome,
}: ThreadViewProps) {
const [chatInput, setChatInput] = useState("");
const plannerThreadId = stream.values?.plannerThreadId;
const [programmerThreadId, setProgrammerThreadId] = useState("");
const plannerThreadId = stream.values?.plannerSession?.threadId;
const plannerRunId = stream.values?.plannerSession?.runId;
const [programmerSession, setProgrammerSession] =
useState<ManagerGraphState["programmerSession"]>();
if (!stream.messages?.length) {
return null;
}
@ -172,14 +175,17 @@ export function ThreadView({
</CardTitle>
</CardHeader>
<CardContent className="space-y-2 p-3 pt-0">
{plannerThreadId && PLANNER_ASSISTANT_ID && (
<ActionsRenderer<PlannerGraphState>
graphId={PLANNER_ASSISTANT_ID}
threadId={plannerThreadId}
setProgrammerThreadId={setProgrammerThreadId}
programmerThreadId={programmerThreadId}
/>
)}
{plannerThreadId &&
plannerRunId &&
PLANNER_ASSISTANT_ID && (
<ActionsRenderer<PlannerGraphState>
graphId={PLANNER_ASSISTANT_ID}
threadId={plannerThreadId}
runId={plannerRunId}
setProgrammerSession={setProgrammerSession}
programmerSession={programmerSession}
/>
)}
</CardContent>
</Card>
</TabsContent>
@ -191,10 +197,11 @@ export function ThreadView({
</CardTitle>
</CardHeader>
<CardContent className="space-y-2 p-3 pt-0">
{programmerThreadId && PROGRAMMER_ASSISTANT_ID && (
{programmerSession && PROGRAMMER_ASSISTANT_ID && (
<ActionsRenderer<PlannerGraphState>
graphId={PROGRAMMER_ASSISTANT_ID}
threadId={programmerThreadId}
threadId={programmerSession.threadId}
runId={programmerSession.runId}
/>
)}
</CardContent>

View file

@ -1,5 +1,5 @@
import { MessagesZodState } from "@langchain/langgraph";
import { TargetRepository, TaskPlan } from "../types.js";
import { TargetRepository, TaskPlan, AgentSession } from "../types.js";
import { z } from "zod";
export const ManagerGraphStateObj = MessagesZodState.extend({
@ -22,13 +22,13 @@ export const ManagerGraphStateObj = MessagesZodState.extend({
*/
taskPlan: z.custom<TaskPlan>(),
/**
* The programmer thread ID
* The programmer session
*/
programmerThreadId: z.string().optional(),
programmerSession: z.custom<AgentSession>().optional(),
/**
* The planner thread ID
* The planner session
*/
plannerThreadId: z.string().optional(),
plannerSession: z.custom<AgentSession>().optional(),
/**
* The branch name to checkout and make changes on.
* Can be user specified, or defaults to `open-swe/<manager-thread-id>

View file

@ -1,7 +1,7 @@
import "@langchain/langgraph/zod";
import { z } from "zod";
import { MessagesZodState } from "@langchain/langgraph";
import { TargetRepository, TaskPlan } from "../types.js";
import { AgentSession, TargetRepository, TaskPlan } from "../types.js";
import { withLangGraph } from "@langchain/langgraph/zod";
export const PlannerGraphStateObj = MessagesZodState.extend({
@ -61,9 +61,9 @@ export const PlannerGraphStateObj = MessagesZodState.extend({
fn: (_state, update) => update,
},
}),
programmerThreadId: withLangGraph(z.custom<string>(), {
programmerSession: withLangGraph(z.custom<AgentSession>(), {
reducer: {
schema: z.custom<string>(),
schema: z.custom<AgentSession>(),
fn: (_state, update) => update,
},
}),

View file

@ -528,3 +528,8 @@ export type GraphConfig = LangGraphRunnableConfig<
assistant_id: string;
}
>;
export interface AgentSession {
threadId: string;
runId: string;
}

View file

@ -2347,7 +2347,7 @@ __metadata:
globals: ^15.14.0
jsonwebtoken: ^9.0.2
katex: ^0.16.21
langgraph-nextjs-api-passthrough: ^0.1.2
langgraph-nextjs-api-passthrough: ^0.1.3
lodash: ^4.17.21
lucide-react: ^0.476.0
next: ^15.2.3
@ -9447,12 +9447,12 @@ __metadata:
languageName: node
linkType: hard
"langgraph-nextjs-api-passthrough@npm:^0.1.2":
version: 0.1.2
resolution: "langgraph-nextjs-api-passthrough@npm:0.1.2"
"langgraph-nextjs-api-passthrough@npm:^0.1.3":
version: 0.1.3
resolution: "langgraph-nextjs-api-passthrough@npm:0.1.3"
peerDependencies:
next: "*"
checksum: 3ed989353b376e8e7b36cf37181838a439246e19cca1e1f8ae3cc27852ce8dcc7ed86ee62acb6c34055c60aea50a248d7f098012e2993f2b2c8e0f650ce71d23
checksum: dc717468620f0902c40e00c7e3a14d279ecb07d9f68ef68bcfa728173bebdab36dd284dbfd11c5d23a554f841709f85c7512f18fa6c2aa75ac73b0fcc4b75976
languageName: node
linkType: hard