diff --git a/libs/langgraph/langgraph/pregel/debug.py b/libs/langgraph/langgraph/pregel/debug.py index 0516c0fd1..53d49f7e1 100644 --- a/libs/langgraph/langgraph/pregel/debug.py +++ b/libs/langgraph/langgraph/pregel/debug.py @@ -186,14 +186,14 @@ def map_debug_checkpoint( "id": t.id, "name": t.name, "error": t.error, - "state": task_states.get(t.id) if task_states else None, + "state": t.state, } if t.error else { "id": t.id, "name": t.name, "interrupts": tuple(asdict(i) for i in t.interrupts), - "state": task_states.get(t.id) if task_states else None, + "state": t.state, } for t in tasks_w_writes(tasks, pending_writes, task_states) ], diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 89088bd1d..7357a0f43 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -6996,13 +6996,13 @@ def test_branch_then( # test stream_mode=debug tool_two = tool_two_graph.compile(checkpointer=checkpointer) thread10 = {"configurable": {"thread_id": "10"}} - + res = [ *tool_two.stream( {"my_key": "value", "market": "DE"}, thread10, stream_mode="debug" ) ] - + assert res == [ { "type": "checkpoint", @@ -7029,7 +7029,14 @@ def test_branch_then( }, "parent_config": None, "next": ["__start__"], - "tasks": [{"id": AnyStr(), "name": "__start__", "interrupts": (), "state": None}], + "tasks": [ + { + "id": AnyStr(), + "name": "__start__", + "interrupts": (), + "state": None, + } + ], }, }, { @@ -7070,7 +7077,9 @@ def test_branch_then( }, }, "next": ["prepare"], - "tasks": [{"id": AnyStr(), "name": "prepare", "interrupts": (), "state": None}], + "tasks": [ + {"id": AnyStr(), "name": "prepare", "interrupts": (), "state": None} + ], }, }, { @@ -7134,7 +7143,14 @@ def test_branch_then( }, }, "next": ["tool_two_slow"], - "tasks": [{"id": AnyStr(), "name": "tool_two_slow", "interrupts": (), "state": None}], + "tasks": [ + { + "id": AnyStr(), + "name": "tool_two_slow", + "interrupts": (), + "state": None, + } + ], }, }, { @@ -7198,7 +7214,9 @@ def test_branch_then( }, }, "next": ["finish"], - "tasks": [{"id": AnyStr(), "name": "finish", "interrupts": (), "state": None}], + "tasks": [ + {"id": AnyStr(), "name": "finish", "interrupts": (), "state": None} + ], }, }, { @@ -11564,3 +11582,75 @@ def test_enum_node_names(): graph = graph.compile() assert graph.invoke({"foo": "hello"}) == {"foo": "hello", "bar": "hello!"} + + +def test_debug_subgraphs(): + class State(TypedDict): + messages: Annotated[list[str], operator.add] + + def node(name): + def _node(state: State): + return {"messages": [f"entered {name} node"]} + + return _node + + grand_parent = StateGraph(State) + parent = StateGraph(State) + child = StateGraph(State) + + child.add_node("c_one", node("c_one")) + child.add_node("c_two", node("c_two")) + child.add_edge(START, "c_one") + child.add_edge("c_one", "c_two") + child.add_edge("c_two", END) + + parent.add_node("p_one", node("p_one")) + parent.add_node("p_two", child.compile()) + parent.add_edge(START, "p_one") + parent.add_edge("p_one", "p_two") + parent.add_edge("p_two", END) + + grand_parent.add_node("gp_one", node("gp_one")) + grand_parent.add_node("gp_two", parent.compile()) + grand_parent.add_edge(START, "gp_one") + grand_parent.add_edge("gp_one", "gp_two") + grand_parent.add_edge("gp_two", END) + + graph = grand_parent.compile(checkpointer=MemorySaver()) + + config = {"configurable": {"thread_id": "1"}} + events = [ + *graph.stream( + {"messages": []}, + config=config, + stream_mode="debug", + ) + ] + + checkpoint_events = list( + reversed([e["payload"] for e in events if e["type"] == "checkpoint"]) + ) + checkpoint_history = list(graph.get_state_history(config)) + + assert len(checkpoint_events) == len(checkpoint_history) + + def normalize_config(config: dict | None) -> dict | None: + if config is None: + return None + return config["configurable"] + + for stream, history in zip(checkpoint_events, checkpoint_history): + assert stream["values"] == history.values + assert stream["next"] == list(history.next) + assert normalize_config(stream["config"]) == normalize_config(history.config) + assert normalize_config(stream["parent_config"]) == normalize_config( + history.parent_config + ) + + assert len(stream["tasks"]) == len(history.tasks) + for stream_task, history_task in zip(stream["tasks"], history.tasks): + assert stream_task["id"] == history_task.id + assert stream_task["name"] == history_task.name + assert stream_task["interrupts"] == history_task.interrupts + assert stream_task.get("error") == history_task.error + assert stream_task.get("state") == history_task.state