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
This commit is contained in:
Caspar Broekhuizen
2025-10-03 09:06:58 -07:00
committed by GitHub
parent 04fb14d3ae
commit b0958115c1
4 changed files with 232 additions and 70 deletions
+80 -44
View File
@@ -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)
+12 -4
View File
@@ -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": [],
},
+12 -4
View File
@@ -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": [],
},
+128 -18
View File
@@ -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": ""},
},
]