diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index fde945709..417e52a55 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -43,7 +43,6 @@ from langgraph.constants import ( MISSING, NS_END, NS_SEP, - SELF, TAG_HIDDEN, ) from langgraph.errors import ( @@ -670,9 +669,9 @@ class StateGraph(Graph): for key, node in self.nodes.items(): compiled.attach_node(key, node) - compiled.attach_branch(START, SELF, CONTROL_BRANCH, with_reader=False) - for key, node in self.nodes.items(): - compiled.attach_branch(key, SELF, CONTROL_BRANCH, with_reader=False) + compiled.nodes[START].writers.append(CONTROL_BRANCH_PATH) + for key in self.nodes: + compiled.nodes[key].writers.append(CONTROL_BRANCH_PATH) for start, end in self.edges: compiled.attach_edge(start, end) @@ -1084,9 +1083,10 @@ def _coerce_state(schema: Type[Any], input: dict[str, Any]) -> dict[str, Any]: return schema(**input) -def _control_branch(value: Any) -> Sequence[Union[str, Send]]: +def _control_branch(value: Any, config: RunnableConfig) -> Sequence[Union[str, Send]]: if isinstance(value, Send): - return [value] + ChannelWrite.do_write(config, (value,)) + return value commands: list[Command] = [] if isinstance(value, Command): commands.append(value) @@ -1094,22 +1094,32 @@ def _control_branch(value: Any) -> Sequence[Union[str, Send]]: for cmd in value: if isinstance(cmd, Command): commands.append(cmd) - rtn: list[Union[str, Send]] = [] + rtn: list[Union[ChannelWriteEntry, Send]] = [] for command in commands: if command.graph == Command.PARENT: raise ParentCommand(command) if isinstance(command.goto, Send): rtn.append(command.goto) elif isinstance(command.goto, str): - rtn.append(command.goto) + rtn.append(ChannelWriteEntry(CHANNEL_BRANCH_TO.format(command.goto), None)) else: - rtn.extend(command.goto) - return rtn + rtn.extend( + go + if isinstance(go, Send) + else ChannelWriteEntry(CHANNEL_BRANCH_TO.format(go), None) + for go in command.goto + ) + if rtn: + ChannelWrite.do_write(config, rtn) + return value -async def _acontrol_branch(value: Any) -> Sequence[Union[str, Send]]: +async def _acontrol_branch( + value: Any, config: RunnableConfig +) -> Sequence[Union[str, Send]]: if isinstance(value, Send): - return [value] + ChannelWrite.do_write(config, (value,)) + return value commands: list[Command] = [] if isinstance(value, Command): commands.append(value) @@ -1117,17 +1127,24 @@ async def _acontrol_branch(value: Any) -> Sequence[Union[str, Send]]: for cmd in value: if isinstance(cmd, Command): commands.append(cmd) - rtn: list[Union[str, Send]] = [] + rtn: list[Union[ChannelWriteEntry, Send]] = [] for command in commands: if command.graph == Command.PARENT: raise ParentCommand(command) if isinstance(command.goto, Send): rtn.append(command.goto) elif isinstance(command.goto, str): - rtn.append(command.goto) + rtn.append(ChannelWriteEntry(CHANNEL_BRANCH_TO.format(command.goto), None)) else: - rtn.extend(command.goto) - return rtn + rtn.extend( + go + if isinstance(go, Send) + else ChannelWriteEntry(CHANNEL_BRANCH_TO.format(go), None) + for go in command.goto + ) + if rtn: + ChannelWrite.do_write(config, rtn) + return value CONTROL_BRANCH_PATH = RunnableCallable( @@ -1136,7 +1153,7 @@ CONTROL_BRANCH_PATH = RunnableCallable( tags=[TAG_HIDDEN], trace=False, recurse=False, - func_accepts_config=False, + func_accepts_config=True, ) CONTROL_BRANCH = Branch(CONTROL_BRANCH_PATH, None)