From b0958115c1887d93bfd25229b307406920038cc4 Mon Sep 17 00:00:00 2001 From: Caspar Broekhuizen Date: Fri, 3 Oct 2025 09:06:58 -0700 Subject: [PATCH] fix(langgraph): task result from stream mode debug / tasks should match format from get_state_history / get_state (#6233) Overview Python port of https://github.com/langchain-ai/langgraphjs/pull/1551 Introduces `map_task_result_writes` to standardize task result format across `get_state_history` and `map_task_result_writes` response structures. Solves https://github.com/langchain-ai/langgraph/issues/6073 --- libs/langgraph/langgraph/pregel/debug.py | 124 +++++++++------ libs/langgraph/tests/test_large_cases.py | 16 +- .../langgraph/tests/test_large_cases_async.py | 16 +- libs/langgraph/tests/test_pregel.py | 146 +++++++++++++++--- 4 files changed, 232 insertions(+), 70 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/debug.py b/libs/langgraph/langgraph/pregel/debug.py index 507fd587a..0de53e120 100644 --- a/libs/langgraph/langgraph/pregel/debug.py +++ b/libs/langgraph/langgraph/pregel/debug.py @@ -40,7 +40,7 @@ class TaskResultPayload(TypedDict): name: str error: str | None interrupts: list[dict] - result: list[tuple[str, Any]] + result: dict[str, Any] class CheckpointTask(TypedDict): @@ -77,6 +77,38 @@ def map_debug_tasks(tasks: Iterable[PregelExecutableTask]) -> Iterator[TaskPaylo } +def is_multiple_channel_write(value: Any) -> bool: + """Return True if the payload already wraps multiple writes from the same channel.""" + return ( + isinstance(value, dict) + and "$writes" in value + and isinstance(value["$writes"], list) + ) + + +def map_task_result_writes(writes: Sequence[tuple[str, Any]]) -> dict[str, Any]: + """Folds task writes into a result dict and aggregates multiple writes to the same channel. + + If the channel contains a single write, we record the write in the result dict as `{channel: write}` + If the channel contains multiple writes, we record the writes in the result dict as `{channel: {'$writes': [write1, write2, ...]}}`""" + + result: dict[str, Any] = {} + for channel, value in writes: + existing = result.get(channel) + + if existing is not None: + channel_writes = ( + existing["$writes"] + if is_multiple_channel_write(existing) + else [existing] + ) + channel_writes.append(value) + result[channel] = {"$writes": channel_writes} + else: + result[channel] = value + return result + + def map_debug_task_results( task_tup: tuple[PregelExecutableTask, Sequence[tuple[str, Any]]], stream_keys: str | Sequence[str], @@ -90,7 +122,9 @@ def map_debug_task_results( "id": task.id, "name": task.name, "error": next((w[1] for w in writes if w[0] == ERROR), None), - "result": [w for w in writes if w[0] in stream_channels_list or w[0] == RETURN], + "result": map_task_result_writes( + [w for w in writes if w[0] in stream_channels_list or w[0] == RETURN] + ), "interrupts": [ asdict(v) for w in writes @@ -196,54 +230,56 @@ def tasks_w_writes( ), MISSING, ) + task_error = next( + (exc for tid, n, exc in pending_writes if tid == task.id and n == ERROR), + None, + ) + task_interrupts = tuple( + v + for tid, n, vv in pending_writes + if tid == task.id and n == INTERRUPT + for v in (vv if isinstance(vv, Sequence) else [vv]) + ) + + task_writes = [ + (chan, val) + for tid, chan, val in pending_writes + if tid == task.id and chan not in (ERROR, INTERRUPT, RETURN) + ] + + if rtn is not MISSING: + task_result = rtn + elif isinstance(output_keys, str): + # unwrap single channel writes to just the write value + filtered_writes = [ + (chan, val) for chan, val in task_writes if chan == output_keys + ] + mapped_writes = map_task_result_writes(filtered_writes) + task_result = mapped_writes.get(str(output_keys)) if mapped_writes else None + else: + if isinstance(output_keys, str): + output_keys = [output_keys] + # map task result writes to the desired output channels + # repeateed writes to the same channel are aggregated into: {'$writes': [write1, write2, ...]} + filtered_writes = [ + (chan, val) for chan, val in task_writes if chan in output_keys + ] + mapped_writes = map_task_result_writes(filtered_writes) + task_result = mapped_writes if filtered_writes else {} + + has_writes = rtn is not MISSING or any( + w[0] == task.id and w[1] not in (ERROR, INTERRUPT) for w in pending_writes + ) + out.append( PregelTask( task.id, task.name, task.path, - next( - ( - exc - for tid, n, exc in pending_writes - if tid == task.id and n == ERROR - ), - None, - ), - tuple( - v - for tid, n, vv in pending_writes - if tid == task.id and n == INTERRUPT - for v in (vv if isinstance(vv, Sequence) else [vv]) - ), + task_error, + task_interrupts, states.get(task.id) if states else None, - ( - rtn - if rtn is not MISSING - else next( - ( - val - for tid, chan, val in pending_writes - if tid == task.id and chan == output_keys - ), - None, - ) - if isinstance(output_keys, str) - else { - chan: val - for tid, chan, val in pending_writes - if tid == task.id - and ( - chan == output_keys - if isinstance(output_keys, str) - else chan in output_keys - ) - } - ) - if any( - w[0] == task.id and w[1] not in (ERROR, INTERRUPT) - for w in pending_writes - ) - else None, + task_result if has_writes else None, ) ) return tuple(out) diff --git a/libs/langgraph/tests/test_large_cases.py b/libs/langgraph/tests/test_large_cases.py index b215e63bb..30ac14fc8 100644 --- a/libs/langgraph/tests/test_large_cases.py +++ b/libs/langgraph/tests/test_large_cases.py @@ -4023,7 +4023,9 @@ def test_in_one_fan_out_out_one_graph_state() -> None: "payload": { "id": AnyStr(), "name": "rewrite_query", - "result": [("query", "query: what is weather in sf")], + "result": { + "query": "query: what is weather in sf", + }, "error": None, "interrupts": [], }, @@ -4071,7 +4073,9 @@ def test_in_one_fan_out_out_one_graph_state() -> None: "payload": { "id": AnyStr(), "name": "retriever_two", - "result": [("docs", ["doc3", "doc4"])], + "result": { + "docs": ["doc3", "doc4"], + }, "error": None, "interrupts": [], }, @@ -4090,7 +4094,9 @@ def test_in_one_fan_out_out_one_graph_state() -> None: "payload": { "id": AnyStr(), "name": "retriever_one", - "result": [("docs", ["doc1", "doc2"])], + "result": { + "docs": ["doc1", "doc2"], + }, "error": None, "interrupts": [], }, @@ -4130,7 +4136,9 @@ def test_in_one_fan_out_out_one_graph_state() -> None: "payload": { "id": AnyStr(), "name": "qa", - "result": [("answer", "doc1,doc2,doc3,doc4")], + "result": { + "answer": "doc1,doc2,doc3,doc4", + }, "error": None, "interrupts": [], }, diff --git a/libs/langgraph/tests/test_large_cases_async.py b/libs/langgraph/tests/test_large_cases_async.py index a8f0b6e19..6754b9935 100644 --- a/libs/langgraph/tests/test_large_cases_async.py +++ b/libs/langgraph/tests/test_large_cases_async.py @@ -2567,7 +2567,9 @@ async def test_in_one_fan_out_out_one_graph_state() -> None: "payload": { "id": AnyStr(), "name": "rewrite_query", - "result": [("query", "query: what is weather in sf")], + "result": { + "query": "query: what is weather in sf", + }, "error": None, "interrupts": [], }, @@ -2615,7 +2617,9 @@ async def test_in_one_fan_out_out_one_graph_state() -> None: "payload": { "id": AnyStr(), "name": "retriever_two", - "result": [("docs", ["doc3", "doc4"])], + "result": { + "docs": ["doc3", "doc4"], + }, "error": None, "interrupts": [], }, @@ -2634,7 +2638,9 @@ async def test_in_one_fan_out_out_one_graph_state() -> None: "payload": { "id": AnyStr(), "name": "retriever_one", - "result": [("docs", ["doc1", "doc2"])], + "result": { + "docs": ["doc1", "doc2"], + }, "error": None, "interrupts": [], }, @@ -2674,7 +2680,9 @@ async def test_in_one_fan_out_out_one_graph_state() -> None: "payload": { "id": AnyStr(), "name": "qa", - "result": [("answer", "doc1,doc2,doc3,doc4")], + "result": { + "answer": "doc1,doc2,doc3,doc4", + }, "error": None, "interrupts": [], }, diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index c894cf720..fcda60881 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -2119,7 +2119,9 @@ def test_in_one_fan_out_state_graph_defer_node( "id": AnyStr(), "name": "rewrite_query", "error": None, - "result": [("query", "query: what is weather in sf")], + "result": { + "query": "query: what is weather in sf", + }, "interrupts": [], }, }, @@ -2153,7 +2155,9 @@ def test_in_one_fan_out_state_graph_defer_node( "id": AnyStr(), "name": "retriever_one", "error": None, - "result": [("docs", ["doc1", "doc2"])], + "result": { + "docs": ["doc1", "doc2"], + }, "interrupts": [], }, }, @@ -2165,7 +2169,9 @@ def test_in_one_fan_out_state_graph_defer_node( "id": AnyStr(), "name": "retriever_two", "error": None, - "result": [("docs", ["doc3", "doc4"])], + "result": { + "docs": ["doc3", "doc4"], + }, "interrupts": [], }, }, @@ -2191,7 +2197,9 @@ def test_in_one_fan_out_state_graph_defer_node( "id": AnyStr(), "name": "analyzer_one", "error": None, - "result": [("query", "analyzed: query: what is weather in sf")], + "result": { + "query": "analyzed: query: what is weather in sf", + }, "interrupts": [], }, }, @@ -2219,7 +2227,9 @@ def test_in_one_fan_out_state_graph_defer_node( "id": AnyStr(), "name": "qa", "error": None, - "result": [("answer", "doc1,doc2,doc3,doc4")], + "result": { + "answer": "doc1,doc2,doc3,doc4", + }, "interrupts": [], }, }, @@ -5539,12 +5549,9 @@ def test_falsy_return_from_task(sync_checkpointer: BaseCheckpointSaver): "id": AnyStr(), "interrupts": [], "name": "falsy_task", - "result": [ - ( - "__return__", - False, - ), - ], + "result": { + "__return__": False, + }, }, "step": 0, "timestamp": AnyStr(), @@ -5561,7 +5568,7 @@ def test_falsy_return_from_task(sync_checkpointer: BaseCheckpointSaver): }, ], "name": "graph", - "result": [], + "result": {}, }, "step": 0, "timestamp": AnyStr(), @@ -5647,12 +5654,9 @@ def test_falsy_return_from_task(sync_checkpointer: BaseCheckpointSaver): "id": AnyStr(), "interrupts": [], "name": "graph", - "result": [ - ( - "__end__", - None, - ), - ], + "result": { + "__end__": None, + }, }, "step": 0, "timestamp": AnyStr(), @@ -8503,3 +8507,109 @@ def test_supersteps_populate_task_results( assert bulk_start_result == ref_start_result == {"num": 1, "text": "one"} assert bulk_double_result == ref_double_result == {"num": 2, "text": "oneone"} + + +def test_multiple_writes_same_channel_from_same_node( + sync_checkpointer: BaseCheckpointSaver, +) -> None: + """Test that a node can write multiple times to the same channel and that writes are ordered, reduced, and reflected in streamed events and state history.""" + + class State(TypedDict): + foo: Annotated[str, lambda a, b: ", ".join([x for x in [a, b] if x])] + + def one(_: State) -> Command: + return Command(update=[("foo", "one.0"), ("foo", "one.1")]) + + def two(_: State) -> State: + return {"foo": "two"} + + graph = ( + StateGraph(State) + .add_node("one", one) + .add_node("two", two) + .add_edge(START, "one") + .add_edge("one", "two") + .add_edge("two", END) + .compile(checkpointer=sync_checkpointer) + ) + + config = {"configurable": {"thread_id": "1"}} + + events = [ + (ns, ev) + for ns, ev in graph.stream( + {"foo": "input"}, config, stream_mode=["updates", "tasks"] + ) + ] + + assert events == [ + ( + "tasks", + { + "id": AnyStr(), + "name": "one", + "input": {"foo": "input"}, + "triggers": ("branch:to:one",), + }, + ), + ("updates", {"one": [{"foo": "one.0"}, {"foo": "one.1"}]}), + ( + "tasks", + { + "id": AnyStr(), + "name": "one", + "error": None, + "result": {"foo": {"$writes": ["one.0", "one.1"]}}, + "interrupts": [], + }, + ), + ( + "tasks", + { + "id": AnyStr(), + "name": "two", + "input": {"foo": "input, one.0, one.1"}, + "triggers": ("branch:to:two",), + }, + ), + ("updates", {"two": {"foo": "two"}}), + ( + "tasks", + { + "id": AnyStr(), + "name": "two", + "error": None, + "result": {"foo": "two"}, + "interrupts": [], + }, + ), + ] + + def map_snapshot(s: StateSnapshot) -> dict: + return { + "tasks": [{"name": t.name, "result": t.result} for t in s.tasks], + "values": s.values, + } + + history = [map_snapshot(s) for s in graph.get_state_history(config)] + + assert history == [ + { + "tasks": [], + "values": {"foo": "input, one.0, one.1, two"}, + }, + { + "tasks": [{"name": "two", "result": {"foo": "two"}}], + "values": {"foo": "input, one.0, one.1"}, + }, + { + "tasks": [ + {"name": "one", "result": {"foo": {"$writes": ["one.0", "one.1"]}}} + ], + "values": {"foo": "input"}, + }, + { + "tasks": [{"name": "__start__", "result": {"foo": "input"}}], + "values": {"foo": ""}, + }, + ]