mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-10 03:37:51 +02:00
Use tuple entry for control branch
This commit is contained in:
@@ -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]]]:
|
||||
|
||||
@@ -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"),
|
||||
],
|
||||
|
||||
Reference in New Issue
Block a user