mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-10 03:37:51 +02:00
Simplify path for control branch attached to every node
- attached to every node to handle command/send return values - used to be a full blown conditional edge, can be simpler by doing all of it in a single function
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user