mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-28 20:45:05 +02:00
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:
@@ -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)
|
||||
|
||||
@@ -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": [],
|
||||
},
|
||||
|
||||
@@ -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": [],
|
||||
},
|
||||
|
||||
@@ -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": ""},
|
||||
},
|
||||
]
|
||||
|
||||
Reference in New Issue
Block a user