Compare commits

...
7 changed files with 1658 additions and 1574 deletions
+1 -1
View File
@@ -13,7 +13,7 @@ env:
jobs: jobs:
build: build:
if: github.ref == 'refs/heads/main' if: github.ref == 'refs/heads/main' || github.ref == 'refs/heads/v0'
runs-on: ubuntu-latest runs-on: ubuntu-latest
outputs: outputs:
+1 -1
View File
@@ -13,7 +13,7 @@ env:
jobs: jobs:
build: build:
if: github.ref == 'refs/heads/main' if: github.ref == 'refs/heads/main' || github.ref == 'refs/heads/v0'
runs-on: ubuntu-latest runs-on: ubuntu-latest
outputs: outputs:
+2 -2
View File
@@ -41,7 +41,7 @@ def run_with_retry(
except ParentCommand as exc: except ParentCommand as exc:
ns: str = config[CONF][CONFIG_KEY_CHECKPOINT_NS] ns: str = config[CONF][CONFIG_KEY_CHECKPOINT_NS]
cmd = exc.args[0] cmd = exc.args[0]
if cmd.graph == ns: if cmd.graph in (ns, task.name):
# this command is for the current graph, handle it # this command is for the current graph, handle it
for w in task.writers: for w in task.writers:
w.invoke(cmd, config) w.invoke(cmd, config)
@@ -137,7 +137,7 @@ async def arun_with_retry(
except ParentCommand as exc: except ParentCommand as exc:
ns: str = config[CONF][CONFIG_KEY_CHECKPOINT_NS] ns: str = config[CONF][CONFIG_KEY_CHECKPOINT_NS]
cmd = exc.args[0] cmd = exc.args[0]
if cmd.graph == ns: if cmd.graph in (ns, task.name):
# this command is for the current graph, handle it # this command is for the current graph, handle it
for w in task.writers: for w in task.writers:
w.invoke(cmd, config) w.invoke(cmd, config)
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project] [project]
name = "langgraph" name = "langgraph"
version = "0.4.7" version = "0.4.8"
description = "Building stateful, multi-actor applications with LLMs" description = "Building stateful, multi-actor applications with LLMs"
authors = [] authors = []
requires-python = ">=3.9" requires-python = ">=3.9"
+45 -2
View File
@@ -5514,8 +5514,11 @@ def test_runnable_passthrough_node_graph() -> None:
assert graph.get_graph(xray=True).to_json() == graph.get_graph(xray=False).to_json() assert graph.get_graph(xray=True).to_json() == graph.get_graph(xray=False).to_json()
@pytest.mark.parametrize("subgraph_persist", [True, False])
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_parent_command(request: pytest.FixtureRequest, checkpointer_name: str) -> None: def test_parent_command(
request: pytest.FixtureRequest, checkpointer_name: str, subgraph_persist: bool
) -> None:
from langchain_core.messages import BaseMessage from langchain_core.messages import BaseMessage
from langchain_core.tools import tool from langchain_core.tools import tool
@@ -5527,7 +5530,7 @@ def test_parent_command(request: pytest.FixtureRequest, checkpointer_name: str)
subgraph_builder = StateGraph(MessagesState) subgraph_builder = StateGraph(MessagesState)
subgraph_builder.add_node("tool", get_user_name) subgraph_builder.add_node("tool", get_user_name)
subgraph_builder.add_edge(START, "tool") subgraph_builder.add_edge(START, "tool")
subgraph = subgraph_builder.compile() subgraph = subgraph_builder.compile(checkpointer=subgraph_persist)
class CustomParentState(TypedDict): class CustomParentState(TypedDict):
messages: Annotated[list[BaseMessage], add_messages] messages: Annotated[list[BaseMessage], add_messages]
@@ -8802,3 +8805,43 @@ def test_imp_exception(
{"my_task": 2}, {"my_task": 2},
{"my_workflow": "done"}, {"my_workflow": "done"},
] ]
@pytest.mark.parametrize("subgraph_persist", [True, False])
def test_parent_command_goto(
sync_checkpointer: BaseCheckpointSaver, subgraph_persist: bool
) -> None:
class State(TypedDict):
dialog_state: Annotated[list[str], operator.add]
def node_a_child(state):
return {"dialog_state": ["a_child_state"]}
def node_b_child(state):
return Command(
graph=Command.PARENT,
goto="node_b_parent",
update={"dialog_state": ["b_child_state"]},
)
sub_builder = StateGraph(State)
sub_builder.add_node(node_a_child)
sub_builder.add_node(node_b_child)
sub_builder.add_edge(START, "node_a_child")
sub_builder.add_edge("node_a_child", "node_b_child")
sub_graph = sub_builder.compile(checkpointer=subgraph_persist)
def node_b_parent(state):
return {"dialog_state": ["node_b_parent"]}
main_builder = StateGraph(State)
main_builder.add_node(node_b_parent)
main_builder.add_edge(START, "subgraph_node")
main_builder.add_node("subgraph_node", sub_graph, destinations=("node_b_parent",))
main_graph = main_builder.compile(sync_checkpointer, name="parent")
config = {"configurable": {"thread_id": 1}}
assert main_graph.invoke(input={"dialog_state": ["init_state"]}, config=config) == {
"dialog_state": ["init_state", "b_child_state", "node_b_parent"]
}
+43 -2
View File
@@ -6772,8 +6772,9 @@ async def test_debug_nested_subgraphs(async_checkpointer: BaseCheckpointSaver):
assert stream_task.get("state") == history_task.state assert stream_task.get("state") == history_task.state
@pytest.mark.parametrize("subgraph_persist", [True, False])
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_parent_command(checkpointer_name: str) -> None: async def test_parent_command(checkpointer_name: str, subgraph_persist: bool) -> None:
from langchain_core.messages import BaseMessage from langchain_core.messages import BaseMessage
from langchain_core.tools import tool from langchain_core.tools import tool
@@ -6785,7 +6786,7 @@ async def test_parent_command(checkpointer_name: str) -> None:
subgraph_builder = StateGraph(MessagesState) subgraph_builder = StateGraph(MessagesState)
subgraph_builder.add_node("tool", get_user_name) subgraph_builder.add_node("tool", get_user_name)
subgraph_builder.add_edge(START, "tool") subgraph_builder.add_edge(START, "tool")
subgraph = subgraph_builder.compile() subgraph = subgraph_builder.compile(checkpointer=subgraph_persist)
class CustomParentState(TypedDict): class CustomParentState(TypedDict):
messages: Annotated[list[BaseMessage], add_messages] messages: Annotated[list[BaseMessage], add_messages]
@@ -9446,3 +9447,43 @@ async def test_imp_exception(
"parent_ids": [], "parent_ids": [],
}, },
] ]
@pytest.mark.parametrize("subgraph_persist", [True, False])
async def test_parent_command_goto(
async_checkpointer: BaseCheckpointSaver, subgraph_persist: bool
) -> None:
class State(TypedDict):
dialog_state: Annotated[list[str], operator.add]
async def node_a_child(state):
return {"dialog_state": ["a_child_state"]}
async def node_b_child(state):
return Command(
graph=Command.PARENT,
goto="node_b_parent",
update={"dialog_state": ["b_child_state"]},
)
sub_builder = StateGraph(State)
sub_builder.add_node(node_a_child)
sub_builder.add_node(node_b_child)
sub_builder.add_edge(START, "node_a_child")
sub_builder.add_edge("node_a_child", "node_b_child")
sub_graph = sub_builder.compile(checkpointer=subgraph_persist)
async def node_b_parent(state):
return {"dialog_state": ["node_b_parent"]}
main_builder = StateGraph(State)
main_builder.add_node(node_b_parent)
main_builder.add_edge(START, "subgraph_node")
main_builder.add_node("subgraph_node", sub_graph, destinations=("node_b_parent",))
main_graph = main_builder.compile(async_checkpointer, name="parent")
config = {"configurable": {"thread_id": 1}}
assert await main_graph.ainvoke(
input={"dialog_state": ["init_state"]}, config=config
) == {"dialog_state": ["init_state", "b_child_state", "node_b_parent"]}
+1565 -1565
View File
File diff suppressed because it is too large Load Diff