Use tuple entry for control branch

This commit is contained in:
Nuno Campos
2025-04-11 09:53:54 -07:00
parent 5a7edead8c
commit 64aa1e6cd8
2 changed files with 16 additions and 63 deletions
+12 -59
View File
@@ -44,6 +44,7 @@ from langgraph.constants import (
NS_END,
NS_SEP,
TAG_HIDDEN,
TASKS,
)
from langgraph.errors import (
ErrorCode,
@@ -78,7 +79,7 @@ from langgraph.store.base import BaseStore
from langgraph.types import All, Checkpointer, Command, RetryPolicy
from langgraph.utils.fields import get_field_default
from langgraph.utils.pydantic import create_model
from langgraph.utils.runnable import RunnableCallable, RunnableLike, coerce_to_runnable
from langgraph.utils.runnable import RunnableLike, coerce_to_runnable
logger = logging.getLogger(__name__)
@@ -669,10 +670,6 @@ class StateGraph(Graph):
for key, node in self.nodes.items():
compiled.attach_node(key, node)
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)
@@ -801,6 +798,7 @@ class CompiledStateGraph(CompiledGraph):
ChannelWriteTupleEntry(
mapper=_get_root if output_keys == ["__root__"] else _get_updates
),
ChannelWriteTupleEntry(mapper=_control_branch),
)
# add node and output channel
@@ -935,7 +933,7 @@ class CompiledStateGraph(CompiledGraph):
if end != END:
self.nodes[end].writers.append(
ChannelWrite(
[ChannelWriteEntry(channel_name, end)], tags=[TAG_HIDDEN]
(ChannelWriteEntry(channel_name, end),), tags=[TAG_HIDDEN]
)
)
@@ -1061,10 +1059,9 @@ def _coerce_state(schema: Type[Any], input: dict[str, Any]) -> dict[str, Any]:
return schema(**input)
def _control_branch(value: Any, config: RunnableConfig) -> Any:
def _control_branch(value: Any) -> Sequence[tuple[str, Any]]:
if isinstance(value, Send):
ChannelWrite.do_write(config, (value,))
return value
return ((TASKS, value),)
commands: list[Command] = []
if isinstance(value, Command):
commands.append(value)
@@ -1072,66 +1069,22 @@ def _control_branch(value: Any, config: RunnableConfig) -> Any:
for cmd in value:
if isinstance(cmd, Command):
commands.append(cmd)
rtn: list[Union[ChannelWriteEntry, Send]] = []
rtn: list[tuple[str, Any]] = []
for command in commands:
if command.graph == Command.PARENT:
raise ParentCommand(command)
if isinstance(command.goto, Send):
rtn.append(command.goto)
rtn.append((TASKS, command.goto))
elif isinstance(command.goto, str):
rtn.append(ChannelWriteEntry(CHANNEL_BRANCH_TO.format(command.goto), None))
rtn.append((CHANNEL_BRANCH_TO.format(command.goto), None))
else:
rtn.extend(
go
(TASKS, go)
if isinstance(go, Send)
else ChannelWriteEntry(CHANNEL_BRANCH_TO.format(go), None)
else (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, config: RunnableConfig) -> Any:
if isinstance(value, Send):
ChannelWrite.do_write(config, (value,))
return value
commands: list[Command] = []
if isinstance(value, Command):
commands.append(value)
elif isinstance(value, (list, tuple)):
for cmd in value:
if isinstance(cmd, Command):
commands.append(cmd)
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(ChannelWriteEntry(CHANNEL_BRANCH_TO.format(command.goto), None))
else:
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(
_control_branch,
_acontrol_branch,
tags=[TAG_HIDDEN],
trace=False,
recurse=False,
set_context=False,
func_accepts_config=True,
)
return rtn
def _get_root(input: Any) -> Optional[Sequence[tuple[str, Any]]]:
+4 -4
View File
@@ -4660,7 +4660,7 @@ def test_root_graph(
content="result for query",
name="search_api",
tool_call_id="tool_call123",
id="00000000-0000-4000-8000-000000000040",
id="00000000-0000-4000-8000-000000000033",
)
]
},
@@ -4683,7 +4683,7 @@ def test_root_graph(
content="result for another",
name="search_api",
tool_call_id="tool_call456",
id="00000000-0000-4000-8000-000000000049",
id="00000000-0000-4000-8000-000000000041",
)
]
},
@@ -5387,7 +5387,7 @@ def test_root_graph(
"__root__": [
HumanMessage(
content="what is weather in sf",
id="00000000-0000-4000-8000-000000000083",
id="00000000-0000-4000-8000-000000000070",
),
AIMessage(
content="",
@@ -5407,7 +5407,7 @@ def test_root_graph(
),
AIMessage(content="answer", id="ai2"),
AIMessage(
content="an extra message", id="00000000-0000-4000-8000-000000000107"
content="an extra message", id="00000000-0000-4000-8000-000000000091"
),
HumanMessage(content="what is weather in la"),
],